making general-purpose metric function
This commit is contained in:
10
src/main.py
10
src/main.py
@@ -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=''):
|
||||
|
||||
Reference in New Issue
Block a user