adding tools to render directly from a policy, updating data generator, adding scratch files

This commit is contained in:
Arec
2021-11-09 06:28:42 -08:00
parent 4c8fb77a91
commit 5799d095d9
7 changed files with 645 additions and 50 deletions

View File

@@ -117,7 +117,7 @@ def load_experts(expert_files):
transitions = rollout.flatten_trajectories(trajectories)
return transitions
def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedIncrementingAgent', path=None, min_timesteps=None, min_episodes=None, video=False, env_args={}, policy_args={}):
def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedRouteIncrementingAgent', path=None, min_timesteps=None, min_episodes=None, video=False, env_args={}, policy_args={}):
"""Rollout and save expert demos.
Usage:
@@ -164,13 +164,13 @@ def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedIncreme
def process_experts(filename:str='expert.pkl',
locs:list=None,
tracks:list=None,
env_class:str='NRasterizedIncrementingAgent',
env_class:str='NRasterizedRouteIncrementingAgent',
env_args:dict={'width':36,'height':36,'m_per_px':2},
expert_class:str='NormalizedIntersimpleExpert',
expert_args:dict={'mu':0.001}):
"""
Process all experts in the Interaction Dataset
For now, using NormalizedIntersimpleExpert with NRasterizedIncrementingAgent environment
For now, using NormalizedIntersimpleExpert with NRasterizedRouteIncrementingAgent environment
Args:
filename (str): name for track file

View File

@@ -1,57 +1,38 @@
import stable_baselines3 as sb3
from intersim.envs.intersimple import NRasterized
import intersim
from src.gail.options import RenderOptions
from tqdm import tqdm
def render_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized):
def render_env(model_name='gail_image_multiagent_nocollision', env='NRasterizedRoute', max_frames=600, options=False,
**env_kwargs):
"""
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
environment (str): 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))
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))
Env = intersim.envs.intersimple.__dict__[env]
print(f'Rendering environment with \'{model_name}\' policy')
if not options:
env = Env(**env_kwargs)
obs = env.reset()
for i in tqdm(range(max_frames)):
action, _states = model.predict(obs)
obs, rewards, done, info = env.step(action)
env.render(mode='post')
if done:
break
else:
env = RenderOptions(Env(**env_kwargs))
with tqdm(total=max_frames) as pbar:
for i, s in enumerate(env.sample_ll(model)):
pbar.update(1)
if s['dones'] or i >= max_frames:
break
env.close(filestr='render/'+model_name)
if __name__ == '__main__':
import fire