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

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

View File

@@ -75,7 +75,7 @@ class CnnDiscriminatorFlatAction(torch.nn.Module):
return sa
def forward(self, state, action):
s = self.cnn(state)
s = self.cnn(state.float())
sa = self._concatenate(s, action)
assert sa.ndim == 2
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
import stable_baselines3
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(
expert_data=expert_data,
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
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__':
# %%
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
trajectories = pickle.load(f)
transitions = rollout.flatten_trajectories(trajectories)
generator = train(transitions, epochs=2, expert_batch_size=2, generator_steps=2)
generator = train(transitions)
generator.save(model_name)

View File

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