fixing flataction discriminator to convert to float beforehand, adding necessary forward calls in expert, adding Fire to video creator from model, and trying full run of options gail with new discrimination model

This commit is contained in:
Arec
2021-10-11 08:47:18 -07:00
parent 01752fac12
commit 70e55327dc
4 changed files with 37 additions and 17 deletions

View File

@@ -79,6 +79,12 @@ class DummyVecEnvPolicy(BasePolicy):
actions = [p[0] for p in predictions] actions = [p[0] for p in predictions]
states = [p[1] for p in predictions] states = [p[1] for p in predictions]
return actions, states return actions, states
def forward(self, *args, **kwargs):
raise NotImplementedError()
def _predict(self, *args, **kwargs):
raise NotImplementedError()
def save_video(env, expert): def save_video(env, expert):
env.reset() env.reset()

View File

@@ -75,7 +75,7 @@ class CnnDiscriminatorFlatAction(torch.nn.Module):
return sa return sa
def forward(self, state, action): def forward(self, state, action):
s = self.cnn(state) s = self.cnn(state.float())
sa = self._concatenate(s, action) sa = self._concatenate(s, action)
assert sa.ndim == 2 assert sa.ndim == 2
return self.decoder(sa).squeeze(1) return self.decoder(sa).squeeze(1)

View File

@@ -1,5 +1,5 @@
# %% # %%
from gail.discriminator import CnnDiscriminator from gail.discriminator import CnnDiscriminator, CnnDiscriminatorFlatAction
from imitation.algorithms import adversarial from imitation.algorithms import adversarial
import stable_baselines3 import stable_baselines3
import torch.utils.data import torch.utils.data
@@ -271,7 +271,8 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di
discriminator = adversarial.GAIL( discriminator = adversarial.GAIL(
expert_data=expert_data, expert_data=expert_data,
expert_batch_size=expert_batch_size, expert_batch_size=expert_batch_size,
discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, discrim_kwargs={'discrim_net': CnnDiscriminatorFlatAction(venv)},
#discrim_kwargs={'discrim_net': CnnDiscriminator(venv)},
venv=venv, # unused venv=venv, # unused
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
) )
@@ -299,10 +300,11 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di
# %% # %%
if __name__ == '__main__': if __name__ == '__main__':
# %% # %%
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f: with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
trajectories = pickle.load(f) trajectories = pickle.load(f)
transitions = rollout.flatten_trajectories(trajectories) transitions = rollout.flatten_trajectories(trajectories)
generator = train(transitions, epochs=2, expert_batch_size=2, generator_steps=2) generator = train(transitions)
generator.save(model_name) generator.save(model_name)

View File

@@ -3,19 +3,31 @@ import stable_baselines3 as sb3
from intersim.envs.intersimple import NRasterized from intersim.envs.intersimple import NRasterized
model_name = 'gail_image_singleagent_nocollision' def render_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized):
model = sb3.PPO.load(model_name) """
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
"""
env = NRasterized(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=51) model = sb3.PPO.load(model_name)
obs = env.reset() env = environment(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=agent)
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) 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__':
import fire
fire.Fire(render_env)