making general-purpose metric function

This commit is contained in:
Arec
2021-07-20 07:05:09 -07:00
parent 226a427436
commit 1a74fa5237
2 changed files with 51 additions and 6 deletions

View File

@@ -4,7 +4,7 @@ import intersim
import numpy as np
import os
opj = os.path.join
from src import InteractionDatasetSingleAgent
from src import InteractionDatasetSingleAgent, metrics
def basestr(**kwargs):
"""
@@ -31,15 +31,13 @@ def main(method='bc', train=False, test=False, loc=0, **kwargs):
os.mkdir(outdir)
filestr = opj(outdir, basestr(**kwargs))
# define transforms
transforms={}
# method-based training
if method=='bc':
from src import bc
policy_class = bc.BehaviorCloningPolicy
load_policy_fn = bc.load_policy
metrics_fn = bc.metrics
train_fn = bc.train
else:
raise NotImplementedError
@@ -60,8 +58,8 @@ def main(method='bc', train=False, test=False, loc=0, **kwargs):
simulate_policy(policy, loc=loc, track=track, filestr=filestr)
# run test metrics
# test_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[4])
# metrics_fn(test_dataset, policy)
test_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[4])
metrics(filestr, test_dataset, policy)
def simulate_policy(policy, loc=0, track=0, filestr=''):