From ffb16cfc31c6d4c7242787406257c5116889ac3d Mon Sep 17 00:00:00 2001 From: Arec Date: Fri, 15 Oct 2021 02:39:49 -0700 Subject: [PATCH] updating function to process all expert data from track files, starting processing options policy from file --- scratch/arec/intersimple/data/expert.py | 4 +-- scratch/arec/intersimple/data/generate.sh | 3 ++- .../intersimple/data/process_all_experts.py | 22 ++++++++++------ .../arec/intersimple/render_env_from_model.py | 25 +++++++++++++++++++ 4 files changed, 43 insertions(+), 11 deletions(-) diff --git a/scratch/arec/intersimple/data/expert.py b/scratch/arec/intersimple/data/expert.py index 3c2986c..587b68d 100644 --- a/scratch/arec/intersimple/data/expert.py +++ b/scratch/arec/intersimple/data/expert.py @@ -111,9 +111,7 @@ def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedIncreme env_args (dict): dictionary of kwargs when instantiating environment class policy_args (dict): dictionary of kwargs when instantiating Expert policy """ - import pdb - pdb.set_trace() - + Env = intersim.envs.intersimple.__dict__[env] Expert = globals()[expert] diff --git a/scratch/arec/intersimple/data/generate.sh b/scratch/arec/intersimple/data/generate.sh index fa75525..90918c9 100755 --- a/scratch/arec/intersimple/data/generate.sh +++ b/scratch/arec/intersimple/data/generate.sh @@ -5,4 +5,5 @@ # python -m expert --env=NRasterizedRandomAgent --min_timesteps=10000 --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentw36h36mppx2.pkl' #python -m expert --env=NRasterized --min_timesteps=200 --env_args='{agent:51,width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl' #python -m expert --env=NRasterized --min_timesteps=3000 --video --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001_NRasterizedRandomAgentw36h36mppx2.pkl' -python -m expert --env=NRasterizedIncrementingAgent --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001_NRasterizedIncrementingAgentw36h36mppx2.pkl' +#python -m expert --env=NRasterizedIncrementingAgent --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001_NRasterizedIncrementingAgentw36h36mppx2.pkl' +python -m process_all_experts --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' diff --git a/scratch/arec/intersimple/data/process_all_experts.py b/scratch/arec/intersimple/data/process_all_experts.py index d4567b7..1499e7b 100644 --- a/scratch/arec/intersimple/data/process_all_experts.py +++ b/scratch/arec/intersimple/data/process_all_experts.py @@ -1,28 +1,34 @@ import tqdm import expert import copy -import sys, os +import os +import intersim +from tqdm import tqdm -def process_all_experts(env_args={}, policy_args={}): +def process_all_experts(filename='expert.pkl',env_args={}, policy_args={}): """ Process all experts in the Interaction Dataset For now, using NormalizedIntersimpleExpert with NRasterizedIncrementingAgent environment Args: + filename (str): name for track file env_args (dict): default environment kwargs policy_args (dict): default policy kwargs """ - - for loc in LOCATIONS: - for track in TRACKS: + I, J = len(intersim.LOCATIONS), intersim.MAX_TRACKS + pbar = tqdm(total=I*J) + for loc in range(I): + for track in range(J): it_env_args = copy.deepcopy(env_args) it_env_args.update({ 'loc':loc, 'track':track, }) - - it_path = 'newpathname' + out_folder = os.path.join(intersim.LOCATIONS[loc], 'track%04i'%(track)) + if not os.path.isdir(out_folder): + os.makedirs(out_folder) + it_path = os.path.join(out_folder,filename) expert.demonstrations( expert='NormalizedIntersimpleExpert', @@ -31,6 +37,8 @@ def process_all_experts(env_args={}, policy_args={}): env_args=it_env_args, policy_args=policy_args, ) + pbar.update(1) + pbar.close() if __name__=='__main__': diff --git a/scratch/arec/intersimple/render_env_from_model.py b/scratch/arec/intersimple/render_env_from_model.py index 0214667..8223c73 100644 --- a/scratch/arec/intersimple/render_env_from_model.py +++ b/scratch/arec/intersimple/render_env_from_model.py @@ -28,6 +28,31 @@ def render_env(model_name='gail_image_multiagent_nocollision', agent=51, environ env.close(filestr='render/'+model_name+'_agent%i'%(agent)) +def render_options_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized): + """ + Render a video from an model, agent, and environment + Args: + model_name (str): name of the model + agent (int): agent to start the video from + environment (gym.Env): gym environment class to render environment on + """ + + model = sb3.PPO.load(model_name) + + env = environment(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=agent) + + obs = env.reset() + i=0 + while True and i < 600: + i+=1 + action, _states = model.predict(obs) + obs, rewards, done, info = env.step(action) + env.render(mode='post') + if done: + break + + env.close(filestr='render/'+model_name+'_agent%i'%(agent)) + if __name__ == '__main__': import fire fire.Fire(render_env) \ No newline at end of file