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

@@ -3,7 +3,7 @@
# tracks:list=None, (default to all tracks) # tracks:list=None, (default to all tracks)
# env_class:str='NRasterizedIncrementingAgent', # env_class:str='NRasterizedIncrementingAgent',
# env_args:dict={width:36,height:36,m_per_px:2}, # env_args:dict={width:36,height:36,m_per_px:2},
# expert_class:str='NormalizedIntersimpleExpert', # expert_class:str='NRasterizedRouteIncrementingAgent',
# expert_args:dict={mu:0.001}): # expert_args:dict={mu:0.001}):
python -m src.data.expert --locs='[DR_USA_Roundabout_FT]' python -m src.data.expert --locs='[DR_USA_Roundabout_FT]' --tracks='[0]'

View File

@@ -0,0 +1 @@
python -m render_options --model_name='gail_options_image_longlong_nocollision' --env='NRasterizedRoute' --options=True --width=36 --height=36 --m_per_px=2 --agent=50 --stop_on_collision=False

View File

@@ -0,0 +1,559 @@
# %%
from gail.discriminator import CnnDiscriminator, CnnDiscriminatorFlatAction
from imitation.algorithms import adversarial
import stable_baselines3
import torch.utils.data
import numpy as np
from intersim.envs.intersimple import NRasterized, NRasterizedRandomAgent
import itertools
from torch.distributions import Categorical
import gym
import torch
import pickle
import imitation.data.rollout as rollout
import tempfile
import pathlib
from imitation.util import logger
from stable_baselines3.common.env_util import make_vec_env
from tqdm import tqdm
import logging
logging.basicConfig(level=logging.DEBUG)
ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is safe fallback
class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy):
"""
Class for high-level options policy (generator)
"""
def __init__(self, observation_space, *args, **kwargs):
super().__init__(observation_space['obs'], *args, **kwargs)
def _prior_distribution(self, s):
"""
Return prior distribution over high-level options (before masking)
Args:
s (torch.tensor): observation
Returns:
values (torch.tensor): values from critic
dist (torch.distributions): prior distribution over actions
"""
latent_pi, latent_vf, latent_sde = self._get_latent(s)
distribution = self._get_action_dist_from_latent(latent_pi, latent_sde)
values = self.value_net(latent_vf)
return values, distribution.distribution
def predict(self, obs):
"""
Will mask invalid states before making action selections
Args:
obs: dict with keys:
obs (torch.tensor): (B,o) true observations
mask (torch.tensor): (B,m) mask over valid actions
Returns:
ch (torch.tensor): (B,a) sampled actions
values (torch.tensor): (B,) predicted value at observation
log_probs (torch.tensor): (B,) log probabilities of selected actions
"""
s, m = obs['obs'], obs['mask']
values, prior = self._prior_distribution(s)
posterior = Categorical(prior.probs * m)
ch = posterior.sample()
return ch, values, posterior.log_prob(ch)
def evaluate_actions(self, obs, ch):
"""
Evaluate particular actions
Args:
obs: dict with keys:
obs (torch.tensor): (B,o) true observations
mask (torch.tensor): (B,m) masks over valid actions
ch (torch.tensor): (B,a) selected actions
Returns:
values (torch.tensor): (B,) predicted value at observation
log_probs (torch.tensor): (B,) log probabilities of selected actions
ent (torch.tensor): (B,) entropy of each distribution over actions
"""
s, m = obs['obs'], obs['mask']
values, prior = self._prior_distribution(s)
posterior = Categorical(prior.probs * m)
return values, posterior.log_prob(ch), posterior.entropy() # additional values used by PPO.train
class OptionsEnv(gym.Wrapper):
"""
Wrap an intersimple environment with an options generator
"""
def __init__(self, env, *args, **kwargs):
"""
Initialize wrapped environment and set high-level action and observation spaces
"""
super().__init__(env, *args, **kwargs)
num_hl_options = len(ALL_OPTIONS)
self.action_space = gym.spaces.Discrete(num_hl_options)
self.observation_space = gym.spaces.Dict({
'obs': env.observation_space,
'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)),
})
def _after_choice(self):
pass
def _after_step(self):
pass
def _transitions(self):
raise NotImplementedError('Use `LLOptions` or `HLOptions` for sampling.')
def sample(self, generator):
"""
yield transitions using a generator
Args:
generator (sb3.PPO)
Yields:
"""
self.done = True
while True:
self.episode_start = False
if self.done:
# reset environment
self.s = self.env.reset()
self.m = available_actions(self.env)
self.done = False
self.episode_start = True
# set the action, the value of the start state, and the logprob of the action
# according to the current environment state and mask
self.ch, self.value, self.log_prob = generator.policy.predict({
'obs': torch.tensor(self.s).unsqueeze(0).to(generator.policy.device),
'mask': torch.tensor(self.m).unsqueeze(0).to(generator.policy.device),
})
# store a float list of actions to take given the option selected in the environment
self.plan = list(map(float, generate_plan(self.env, self.ch)))
# run whatever _after_choice might dictate in a child class
self._after_choice()
# some checks
assert not self.done
assert self.plan
assert feasible(self.env, self.plan, self.ch)
# execute the option so long as the episode isn't complete and the plan is still feasible
while not self.done and self.plan and feasible(self.env, self.plan, self.ch):
# pop first action
self.a, self.plan = self.plan[0], self.plan[1:]
# normalize action ??
self.a = self.env._normalize(self.a)
# step through environment
self.nexts, _, self.done, _ = self.env.step(self.a)
self.nextm = available_actions(self.env)
# run whatever _after_step might dictate in child class
self._after_step()
# update state and mask to current
self.s = self.nexts
self.m = self.nextm
# transitions yielded from self._transitions() functions specied in child classes
yield from self._transitions()
### NOTE: only yields after a full option has been executed / exited
class LLOptions(OptionsEnv):
"""Sample low-level (state, action) tuples for discriminator training."""
def __init__(self, *args, **kwargs):
"""
LLOption uses the true LL observations
"""
super().__init__(*args, **kwargs)
# overwrite observation space to just output obs directly
self.observation_space = self.observation_space['obs']
def _after_choice(self):
"""
After each option choice, initialize/reset the transition buffer
"""
self._transition_buffer = []
def _after_step(self):
"""
After each ll action, append s, s', a, done to transition buffer
"""
self._transition_buffer.append({
'obs': self.s,
'next_obs': self.nexts,
'acts': np.array((self.a,)),
'dones': np.array(self.done),
})
def _transitions(self):
"""
Yield from the transition buffer
"""
yield from self._transition_buffer
def sample_ll(self, policy):
"""
Args:
policy
Returns:
gen: iterable which samples low-level transitions from the environment
"""
return self.sample(policy)
class HLOptions(OptionsEnv):
"""Sample high-level (state, action, reward) tuples for generator training."""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def _after_choice(self):
"""
After an option selection, initialize total reward and number of steps
"""
self.r = 0
self.steps = 0
def _after_step(self):
"""
After each low-level action, add the discounted discriminated reward score (given a discriminator)
"""
self.r += self.discount**self.steps * self.discriminator.discrim_net.reward_train(
state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()),
action=torch.tensor([[self.a]]).to(self.discriminator.discrim_net.device()),
next_state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused
done=torch.tensor(self.done).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused
)
self.steps += 1
def _transitions(self):
"""
Yield a single dictionary per high-level selected action
Fields:
obs: high-level state and mask at selection
action: chosen high-level action
reward: accumulated option reward
episode_start: whether the action was chosen at the episode start
value: the value estimate from the starting state
log_prob: the log_prob of the selected action from the starting state
done: whether the episode has ended
"""
yield {
'obs': {'obs': self.s, 'mask': self.m},
'action': self.ch,
'reward': self.r.detach(),
'episode_start': self.episode_start,
'value': self.value.detach(),
'log_prob': self.log_prob.detach(),
'done': self.done,
}
def sample_hl(self, policy, discriminator):
"""
Args:
policy
discriminator: function with which to score rewards
Returns:
gen: iterable which samples high-level transitions from the environment
"""
self.discriminator = discriminator
return self.sample(policy)
class RenderOptions(LLOptions):
def _after_step(self):
"""
Render the environment after each low-level step
"""
super()._after_step()
self.env.render()
def close(self, *args, **kwargs):
"""
On 'close', close the environment
"""
self.env.close(*args, **kwargs)
def available_actions(env):
"""Return mask of available actions given current `env` state."""
valid = np.array([feasible(env, generate_plan(env, i), i) for i in range(len(ALL_OPTIONS))])
return valid
def target_velocity_plan(current_v: float, target_v: float, t: int, dt: float):
"""Smoothly target a velocity in a given number of steps"""
# for now, constant acceleration
a = (target_v - current_v) / (t * dt)
return a*np.ones((t,))
def generate_plan(env, i):
"""Generate input profile for high-level action `i`."""
assert i < len(ALL_OPTIONS), "Invalid option index {i}"
target_v, t = ALL_OPTIONS[i]
current_v = env._env.state[env._agent, 1].item() # extract from env
plan = target_velocity_plan(current_v, target_v, t, env._env._dt)
assert len(plan) == t, "incorrect plan length"
return plan
def check_future_collisions_fast(env, actions):
"""Checks whether `env._agent` would collide with other agents assuming `actions` as input.
Vehicles are (over-)approximated by single circles.
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free
"""
B, (T, nv, _) = len(actions), actions[0].shape
states = torch.stack(env._env.propagate_action_profile(actions), axis=0)
assert states.shape == (B, T, nv, 5)
distance = ((states[:, :, :, :2] - states[:, :, env._agent:env._agent+1, :2])**2).sum(-1).sqrt()
distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents
distance[:, :, env._agent] = np.inf # cannot collide with itself
assert distance.shape == (B, T, nv)
radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2
min_distance = radius[env._agent] + radius
min_distance = min_distance.unsqueeze(0).unsqueeze(0)
assert min_distance.shape == (1, 1, nv)
return (distance > min_distance).all(-1).all(-1)
def check_future_collisions_circles(env, actions, n_circles:int=2):
"""Checks whether `env._agent` would collide with other agents assuming `actions` as input.
Vehicles are (over-)approximated by multiple circles.
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free
"""
assert n_circles >= 2
B, (T, nv, _) = len(actions), actions[0].shape
states = torch.stack(env._env.propagate_action_profile(actions), axis=0)
assert states.shape == (B, T, nv, 5)
centers = states[:, :, :, :2]
psi = states[:, :, :, 3]
lon = torch.stack([psi.cos(), psi.sin()],dim=-1) # (B, T, nv, 2)
# offset between [-env._env.lengths+env._env.widths/2, env._env.lengths/2-env._env.widths/2]
back = (-env._env._lengths/2+env._env._widths/2).unsqueeze(-1) # (nv, 1)
length = (env._env._lengths-env._env._widths).unsqueeze(-1) # (nv, 1)
diff_d = back + length*(torch.arange(n_circles)/(n_circles-1)).unsqueeze(0) # (nv, n_circles)
assert diff_d.shape == (nv, n_circles)
offsets = diff_d[None, None, :, :, None] * lon[:, :, :, None, :]
assert offsets.shape == (B, T, nv, n_circles, 2)
expanded_centers=centers.unsqueeze(-2) + offsets #(B, T, nv, n_circles, 2)
assert expanded_centers.shape == (B, T, nv, n_circles, 2)
agent_centers = expanded_centers[:,:,env._agent:env._agent+1,:,:] #(B, T, 1, n_circles, 2)
ds = expanded_centers.reshape((B, T, nv*n_circles, 1, 2)) - agent_centers #(B, T, nv*nc,1, 2) - (B, T, 1, nc, 2) = (B, T, nv*nc, nc, 2)
distance = (ds**2).sum(-1).sqrt().reshape((B, T, nv, n_circles, n_circles)) # (B, T, nv, nc, nc)
distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents
distance[:, :, env._agent] = np.inf # cannot collide with itself
assert distance.shape == (B, T, nv, n_circles, n_circles)
radius = env._env._widths*np.sqrt(2) / 2
min_distance = radius[env._agent] + radius
min_distance = min_distance[None, None, :, None, None]
assert min_distance.shape == (1, 1, nv, 1, 1)
return (distance > min_distance).all(-1).all(-1).all(-1).all(-1)
def feasible(env, plan, ch):
"""Check if input profile is feasible given current `env` state. Action `ch=0` is safe fallback."""
# zero pad plan - Take (T,) np plan and convert it to (T, nv, 1) torch.Tensor
full_plan = torch.zeros(len(plan), env._env._nv, 1)
full_plan[:, env._agent, 0] = torch.tensor(plan)
# valid = check_future_collisions_fast(env, [full_plan]) # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor
valid = check_future_collisions_circles(env, [full_plan])
return ch == 0 or valid.item()
def flatten_transitions(transitions):
return {
'obs': np.stack(list(t['obs'] for t in transitions), axis=0),
'next_obs': np.stack(list(t['next_obs'] for t in transitions), axis=0),
'acts': np.stack(list(t['acts'] for t in transitions), axis=0),
'dones': np.stack(list(t['dones'] for t in transitions), axis=0),
}
def train_discriminator(env, generator, discriminator, num_samples):
transitions = list(itertools.islice(env.sample_ll(generator), num_samples))
generator_samples = flatten_transitions(transitions)
discriminator.train_disc(gen_samples=generator_samples)
def train_generator(env, generator, discriminator, num_samples):
generator_samples = list(itertools.islice(env.sample_hl(generator, discriminator), num_samples+1))
generator.rollout_buffer.reset()
for s in generator_samples[:-1]:
generator.rollout_buffer.add(
obs=s['obs'],
action=s['action'].cpu(),
reward=s['reward'].cpu(),
episode_start=s['episode_start'],
value=s['value'],
log_prob=s['log_prob'],
)
generator.rollout_buffer.compute_returns_and_advantage(
last_values=generator_samples[-1]['value'],
dones=generator_samples[-1]['done'],
)
generator.train()
def train(expert_data, env_class=NRasterizedRandomAgent, env_settings={}, epochs=10, discrim_batch_size=32, generator_steps=2048, discount=0.99):
"""
Args:
expert_data: list of transitions
env_class: environment class
env_settings: environment settings
epochs: number of epochs to train for
discrim_batch_size: discriminator batch size
generator_steps: number of steps taken in generator
discount: discount factor
Returns:
generator (stable_baselines3.PPO): options policy
"""
env = env_class(**env_settings)
env.discount = discount
tempdir = tempfile.TemporaryDirectory(prefix="quickstart")
tempdir_path = pathlib.Path(tempdir.name)
logger.configure(tempdir_path / "GAIL/")
print(f"All Tensorboards and logging are being written inside {tempdir_path}/.")
venv = make_vec_env(env_class, n_envs=1, env_kwargs=env_settings)
discriminator = adversarial.GAIL(
expert_data=expert_data,
expert_batch_size=discrim_batch_size,
discrim_kwargs={'discrim_net': CnnDiscriminatorFlatAction(venv)},
#discrim_kwargs={'discrim_net': CnnDiscriminator(venv)},
venv=venv, # unused
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
)
generator = stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env),
verbose=1,
n_steps=generator_steps,
)
# PPO.train requires logger as set up in
# PPO._setup_learn (called by PPO.learn)
generator._logger = stable_baselines3.common.utils.configure_logger(
generator.verbose,
generator.tensorboard_log,
)
for _ in tqdm(range(epochs)):
train_discriminator(LLOptions(env), generator, discriminator, num_samples=discrim_batch_size)
train_generator(HLOptions(env), generator, discriminator, num_samples=generator_steps)
return generator
# %%
if __name__ == '__main__':
# %%
model_name = 'gail_options_image'
env_class = NRasterizedRandomAgent
env_settings = {'width': 36, 'height': 36, 'm_per_px': 2}
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedIncrementingAgentw36h36mppx2.pkl", "rb") as f:
trajectories = pickle.load(f)
#import pdb
#pdb.set_trace()
transitions = rollout.flatten_trajectories(trajectories)
generator = train(
transitions,
env_class=env_class,
env_settings=env_settings,
epochs=2,
discrim_batch_size=32,
generator_steps=2048,
discount=0.99
)
generator.save(model_name) # save ppo sb3 generator class
# %%
model = stable_baselines3.PPO.load(model_name) # not actually used
env = RenderOptions(NRasterizedRandomAgent(**env_settings))
for s in env.sample_ll(generator):
if s['dones']:
break
env.close(filestr='render/'+model_name)
# %% Tests
def test_ll_expert_data():
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
expert_trajectories = pickle.load(f)
expert_transitions = rollout.flatten_trajectories(expert_trajectories)
env = LLOptions(NRasterized(agent=51, width=36, height=36, m_per_px=2))
gen_transitions = list(itertools.islice(env.sample_ll(
policy=stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env),
verbose=1,
)
), 10))
gen_transitions = flatten_transitions(gen_transitions)
assert expert_transitions[:10].obs.shape == gen_transitions['obs'].shape
assert expert_transitions[:10].next_obs.shape == gen_transitions['next_obs'].shape
assert expert_transitions[:10].acts.shape == gen_transitions['acts'].shape
assert expert_transitions[:10].dones.shape == gen_transitions['dones'].shape
def test_ll_states():
env = NRasterized()
policy = stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env),
verbose=1,
)
llenv = LLOptions(env)
transitions = list(itertools.islice(llenv.sample_ll(policy=policy), 100))
env2 = NRasterized()
s2 = env2.reset()
for i, t in enumerate(transitions):
assert i == 0 or np.array_equal(t['obs'], transitions[i-1]['next_obs'])
assert np.array_equal(t['obs'], s2)
assert t['acts'].shape == (1,)
nexts2, _, done2, _ = env2.step(t['acts'])
assert np.array_equal(t['next_obs'], nexts2)
assert np.array_equal(t['dones'], done2)
if done2:
break
s2 = nexts2
def test_hl_transitions():
pass

View File

@@ -0,0 +1,46 @@
TRAINING
---------
1. Load pre-trained massive set of transitions
-- For all roundabouts
-- For all tracks
-- For all vehicles
-- For all valid timesteps
-- Rasterized state (incl. path), action
2. HGAIL
-- For each epoch
-- INSTANTIATE A NEW ENVIRONMENT (Roundabout + Track) w/ randomized agent, from set of all expert environments
-- Train discriminator off training data + yielded low-level transitions
-- Train generator off yielded high-level transitions + summed low-level discriminator rewards
TESTING
----------
1. Save average vehicle velocities for all expert vehicles (loop roundabout + track + vehicle, average over time)
2. Run test suite for: expert, BC, GAIL, RAIL, HGAIL, (and hopefully HRAIL)
-- For all roundabouts, tracks
-- Get expert velocities for track
-- Simulate incrementing agent environment (e.g. on reset, agent +=1)
-- Store low-level true joint states, actions, and controlled vehicle index
-- Per-vehicle statistics (v_all, v_mean, v_shortfall, a_all, jerk_all, n_collisions, T)
-- Aggregate statistics + joint
Problems
-----------
Should train without stopping for collisions, however when doing so, end up with policy that always takes decelerate option
-- It seems safe at the start of each vehicles sim, but actually it isn't since a car will spawn and hit it
Solutions:
-- Hold cars from spawning if their spawn location is full
-- Start simulations a few seconds later (after cars clear their spawn places)
Save massive set of transition raw states beforehand (1 from training, but with raw states)
# -- For all roundabouts, tracks
# -- For all vehicles, steps
# -- Raw vehicle state, action

View File

@@ -0,0 +1,8 @@
import sys
sys.path.append('../../../')
from src.util import render_env
if __name__=='__main__':
import fire
fire.Fire(render_env)

View File

@@ -117,7 +117,7 @@ def load_experts(expert_files):
transitions = rollout.flatten_trajectories(trajectories) transitions = rollout.flatten_trajectories(trajectories)
return transitions 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. """Rollout and save expert demos.
Usage: Usage:
@@ -164,13 +164,13 @@ def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedIncreme
def process_experts(filename:str='expert.pkl', def process_experts(filename:str='expert.pkl',
locs:list=None, locs:list=None,
tracks:list=None, tracks:list=None,
env_class:str='NRasterizedIncrementingAgent', env_class:str='NRasterizedRouteIncrementingAgent',
env_args:dict={'width':36,'height':36,'m_per_px':2}, env_args:dict={'width':36,'height':36,'m_per_px':2},
expert_class:str='NormalizedIntersimpleExpert', expert_class:str='NormalizedIntersimpleExpert',
expert_args:dict={'mu':0.001}): expert_args:dict={'mu':0.001}):
""" """
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 NRasterizedRouteIncrementingAgent environment
Args: Args:
filename (str): name for track file filename (str): name for track file

View File

@@ -1,57 +1,38 @@
import stable_baselines3 as sb3 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', env='NRasterizedRoute', max_frames=600, options=False,
def render_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized): **env_kwargs):
""" """
Render a video from an model, agent, and environment Render a video from an model, agent, and environment
Args: Args:
model_name (str): name of the model model_name (str): name of the model
agent (int): agent to start the video from environment (str): gym environment class to render environment on
environment (gym.Env): gym environment class to render environment on
""" """
model = sb3.PPO.load(model_name) model = sb3.PPO.load(model_name)
Env = intersim.envs.intersimple.__dict__[env]
env = environment(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=agent) print(f'Rendering environment with \'{model_name}\' policy')
if not options:
obs = env.reset() env = Env(**env_kwargs)
i=0 obs = env.reset()
while True and i < 600: for i in tqdm(range(max_frames)):
i+=1 action, _states = model.predict(obs)
action, _states = model.predict(obs) obs, rewards, done, info = env.step(action)
obs, rewards, done, info = env.step(action) env.render(mode='post')
env.render(mode='post') if done:
if done: break
break else:
env = RenderOptions(Env(**env_kwargs))
env.close(filestr='render/'+model_name+'_agent%i'%(agent)) with tqdm(total=max_frames) as pbar:
for i, s in enumerate(env.sample_ll(model)):
def render_options_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized): pbar.update(1)
""" if s['dones'] or i >= max_frames:
Render a video from an model, agent, and environment break
Args: env.close(filestr='render/'+model_name)
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