diff --git a/scratch/arec/intersimple/commands.txt b/scratch/arec/intersimple/commands.txt index 1e1285a..b706521 100644 --- a/scratch/arec/intersimple/commands.txt +++ b/scratch/arec/intersimple/commands.txt @@ -1 +1 @@ -python -m render_options --model_name='gail_options_image_longlong_nocollision' --env='NRasterizedRoute' --options=True --width=36 --height=36 --m_per_px=2 --agent=50 --stop_on_collision=False \ No newline at end of file +python -m render_options --model_name='gail_options_image_mid_wcollision' --env='NRasterizedRoute' --options=True --width=36 --height=36 --m_per_px=2 --agent=50 --stop_on_collision=False \ No newline at end of file diff --git a/scratch/arec/intersimple/plan.txt b/scratch/arec/intersimple/plan.txt index 8425d4b..dac52b6 100644 --- a/scratch/arec/intersimple/plan.txt +++ b/scratch/arec/intersimple/plan.txt @@ -36,8 +36,12 @@ Should train without stopping for collisions, however when doing so, end up with -- It seems safe at the start of each vehicles sim, but actually it isn't since a car will spawn and hit it Solutions: -- Hold cars from spawning if their spawn location is full - -- Start simulations a few seconds later (after cars clear their spawn places) + -- Start simulations a few seconds later (after cars clear their spawn places) <- Preferred +Test could run indefinitely if stop_on_collision is off +Solution: + -- Set maximum episode length in intersimple + diff --git a/scratch/arec/intersimple/test_model.py b/scratch/arec/intersimple/test_model.py new file mode 100644 index 0000000..e4a5758 --- /dev/null +++ b/scratch/arec/intersimple/test_model.py @@ -0,0 +1,135 @@ +from tqdm import tqdm +from copy import deepcopy + +ALL_OPTIONS = + +def load_model(model_path:str, method:str): + """ + Load a model given a path and the method + + Args: + model_path (str): the path to the model + method (str): the method for the model + Returns: + model: the action model + is_heir (bool): whether the method is heirarchial + """ + model = None + is_heir = False + if method == 'expert': + pass + elif method == 'bc': + raise NotImplementedError + elif method == 'gail': + raise NotImplementedError + elif method == 'rail': + raise NotImplementedError + elif method == 'hgail': + is_heir = True + raise NotImplementedError + elif method == 'hrail': + is_heir = True + raise NotImplementedError + else: + raise NotImplementedError + return model, is_heir + +def load_expert_states(roundabout, track): + """ + Load expert states from roundabout/track info + Args: + roundabout (str): roundabout name + track (str): track id + Returns: + expert_states (torch.tensor): (nv, T, 5) expert states for track file + """ + pass + +def test_model( + locations=[], + model_name='gail_image_multiagent_nocollision', + env='NRasterizedRouteIncrementingAgent', + method='expert', + options_list=ALL_OPTIONS, + **env_kwargs): + """ + Test a particular model at different locations/tracks + + Args: + locations (list of tuples): list of (roundabout, track) pairs + model_name (str): name of model to test + env (str): environment class + method (str): method (expert, bc, gail, rail, hgail, hrail) + options_list (list): list of options + """ + + # load policy + policy, is_heir = load_model(model_name, method) + + # iterate through vehicles + all_vehicle_infos = [] + for i, location in tqdm(enumerate(locations)): + + # load expert states + expert_states = load_expert_states(roundabout, track) + + # add roundabout and track to environent + roundabout, track = location + it_env_kwargs = deepcopy(env_kwargs) + it_env_kwargs.update({}) + + # load expert states and get average velocities + expert_states = load_expert_states(roundabout, track) + expert_vavg = torch.nanmean(expert_states[:,:,3], dim=-1) + + # initialize environment + if not is_heir: + pass + else: + pass + s = env.reset() + + # Iterate through every vehicle and time + vehicle_infos, done = [], False + for iv in range(env.nv): + v_number = env.agent + i_vehicle_infos = {'s':[], 'a':[], 'it':[]} + while not done: + a = policy(s) + sp, r, done, info = env.step(a) + i_vehicle_infos['s'].append(env._env.state) # FIX + i_vehicle_infos['a'].append(a) + i_vehicle_infos['it'].append(env._env.it) # FIX + i_vehicle_info.update({ + 'vehicle_id': env.agent, + 'n_steps': len(i_vehicle_infos['a']), + 'T': len(i_vehicle_infos['a'])*env._env.dt, # FIX + 'n_collisions': collision.check(i_vehicle_infos['s'], env._env.lengths. env._env.widths), # FIX + 'expert_vavg': expert_vavg[env.agent] + }) + vehicle_infos.append(i_vehicle_info) + env.reset() + + all_vehicle_infos.append({ + 'loc': location, + 'track': track, + 'stats': vehicle_infos + }) + env.close() + + # print and save model-specific metrics + outfolder = 'test_metrics' + print_and_save(all_vehicle_infos, method, model, outfolder) + +def print_and_save(stats, method, model, outfolder): + """ + Print and save stats + """ + pass + +def load_compare(): + pass + +if __name__=='__main__': + import fire + fire.Fire() \ No newline at end of file