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:
@@ -80,6 +80,12 @@ class DummyVecEnvPolicy(BasePolicy):
|
|||||||
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()
|
||||||
env.render()
|
env.render()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -3,14 +3,22 @@ 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:
|
obs = env.reset()
|
||||||
|
i=0
|
||||||
|
while True and i < 600:
|
||||||
i+=1
|
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)
|
||||||
@@ -18,4 +26,8 @@ while True and i < 600:
|
|||||||
if done:
|
if done:
|
||||||
break
|
break
|
||||||
|
|
||||||
env.close(filestr='render/'+model_name)
|
env.close(filestr='render/'+model_name+'_agent%i'%(agent))
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
import fire
|
||||||
|
fire.Fire(render_env)
|
||||||
Reference in New Issue
Block a user