updating function to process all expert data from track files, starting processing options policy from file

This commit is contained in:
Arec
2021-10-15 02:39:49 -07:00
parent 415d607418
commit ffb16cfc31
4 changed files with 43 additions and 11 deletions

View File

@@ -111,8 +111,6 @@ def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedIncreme
env_args (dict): dictionary of kwargs when instantiating environment class env_args (dict): dictionary of kwargs when instantiating environment class
policy_args (dict): dictionary of kwargs when instantiating Expert policy policy_args (dict): dictionary of kwargs when instantiating Expert policy
""" """
import pdb
pdb.set_trace()
Env = intersim.envs.intersimple.__dict__[env] Env = intersim.envs.intersimple.__dict__[env]
Expert = globals()[expert] Expert = globals()[expert]

View File

@@ -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=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=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=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}'

View File

@@ -1,28 +1,34 @@
import tqdm import tqdm
import expert import expert
import copy 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 Process all experts in the Interaction Dataset
For now, using NormalizedIntersimpleExpert with NRasterizedIncrementingAgent environment For now, using NormalizedIntersimpleExpert with NRasterizedIncrementingAgent environment
Args: Args:
filename (str): name for track file
env_args (dict): default environment kwargs env_args (dict): default environment kwargs
policy_args (dict): default policy kwargs policy_args (dict): default policy kwargs
""" """
I, J = len(intersim.LOCATIONS), intersim.MAX_TRACKS
for loc in LOCATIONS: pbar = tqdm(total=I*J)
for track in TRACKS: for loc in range(I):
for track in range(J):
it_env_args = copy.deepcopy(env_args) it_env_args = copy.deepcopy(env_args)
it_env_args.update({ it_env_args.update({
'loc':loc, 'loc':loc,
'track':track, 'track':track,
}) })
out_folder = os.path.join(intersim.LOCATIONS[loc], 'track%04i'%(track))
it_path = 'newpathname' if not os.path.isdir(out_folder):
os.makedirs(out_folder)
it_path = os.path.join(out_folder,filename)
expert.demonstrations( expert.demonstrations(
expert='NormalizedIntersimpleExpert', expert='NormalizedIntersimpleExpert',
@@ -31,6 +37,8 @@ def process_all_experts(env_args={}, policy_args={}):
env_args=it_env_args, env_args=it_env_args,
policy_args=policy_args, policy_args=policy_args,
) )
pbar.update(1)
pbar.close()
if __name__=='__main__': if __name__=='__main__':

View File

@@ -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)) 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__': if __name__ == '__main__':
import fire import fire
fire.Fire(render_env) fire.Fire(render_env)