updating function to process all expert data from track files, starting processing options policy from file
This commit is contained in:
@@ -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]
|
||||||
|
|||||||
@@ -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}'
|
||||||
|
|||||||
@@ -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__':
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user