adding tools to render directly from a policy, updating data generator, adding scratch files
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user