adding discriminators to main folder, utilities to render a video from a saved model

This commit is contained in:
Arec
2021-10-21 05:33:05 -07:00
parent da1fb11269
commit 45a99978e4
5 changed files with 206 additions and 0 deletions

View File

@@ -0,0 +1 @@
from src.discriminator.discriminator import *

View File

@@ -0,0 +1,101 @@
import torch
# imitation.rewards.discrim_nets.DiscrimNetGAIL is composed of self.discriminator (nn.Module),
# which gets called with inputs (state, action) when needed.
class CnnDiscriminator(torch.nn.Module):
"""ConvNet similar to stable_baselines3.common.policies.ActorCriticCnnPolicy."""
def __init__(self, env):
super().__init__()
obs_channels, _, _ = env.observation_space.shape
(action_size,) = env.action_space.shape
in_channels = obs_channels + action_size
self.cnn = torch.nn.Sequential(
torch.nn.Conv2d(in_channels, 32, kernel_size=(8, 8), stride=(4, 4)), # 5+1 -> 32
torch.nn.ReLU(),
torch.nn.Conv2d(32, 64, kernel_size=(4, 4), stride=(2, 2)), # 32 -> 64
torch.nn.ReLU(),
torch.nn.Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1)), # 64 -> 64
torch.nn.ReLU(),
torch.nn.Flatten(start_dim=1, end_dim=-1),
torch.nn.LazyLinear(512), # 28224 -> 512
torch.nn.ReLU(),
torch.nn.LazyLinear(1), # 512 -> 1
)
@staticmethod
def _concatenate(state, action):
b, _, h, w = state.shape
_, a = action.shape
act = action.unsqueeze(-1).unsqueeze(-1).expand((b, a, h, w))
sa = torch.cat((state, act), -3)
return sa
def forward(self, state, action):
sa = self._concatenate(state, action)
assert sa.ndim == 4
return self.cnn(sa).squeeze(1)
class CnnDiscriminatorFlatAction(torch.nn.Module):
"""ConvNet similar to stable_baselines3.common.policies.ActorCriticCnnPolicy."""
def __init__(self, env):
super().__init__()
obs_channels, _, _ = env.observation_space.shape
(action_size,) = env.action_space.shape
in_channels = obs_channels
self.cnn = torch.nn.Sequential(
torch.nn.Conv2d(in_channels, 32, kernel_size=(8, 8), stride=(4, 4)), # in_channels -> 32
torch.nn.ReLU(),
torch.nn.Conv2d(32, 64, kernel_size=(4, 4), stride=(2, 2)), # 32 -> 64
torch.nn.ReLU(),
torch.nn.Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1)), # 64 -> 64
torch.nn.ReLU(),
torch.nn.Flatten(start_dim=1, end_dim=-1),
torch.nn.LazyLinear(128), # 28224 -> 128
)
self.decoder = torch.nn.Sequential(
torch.nn.LazyLinear(64), #128 + 2 -> 64
torch.nn.ReLU(),
torch.nn.LazyLinear(64), #64 -> 64
torch.nn.ReLU(),
torch.nn.LazyLinear(1) #64 -> 1
)
@staticmethod
def _concatenate(state, action):
b, s= state.shape
b, a = action.shape
sa = torch.cat((state, action), -1)
return sa
def forward(self, state, action):
s = self.cnn(state.float())
sa = self._concatenate(s, action)
assert sa.ndim == 2
return self.decoder(sa).squeeze(1)
class MlpDiscriminator(torch.nn.Module):
"""MLP similar to stable_baselines3.common.policies.ActorCriticPolicy."""
def __init__(self, env=None):
super().__init__()
self.flatten = torch.nn.Flatten(start_dim=1, end_dim=-1)
self.mlp = torch.nn.Sequential(
torch.nn.LazyLinear(64), # 42 -> 64
torch.nn.Tanh(),
torch.nn.LazyLinear(64), # 64 -> 64
torch.nn.Tanh(),
torch.nn.LazyLinear(1), # 64 -> 1
)
def forward(self, state, action):
flat = self.flatten(state)
sa = torch.cat((action, flat), -1)
assert sa.ndim == 2
return self.mlp(sa).squeeze(1)

View File

@@ -0,0 +1 @@
from src.util.render_env import *

58
src/util/render_env.py Normal file
View File

@@ -0,0 +1,58 @@
import stable_baselines3 as sb3
from intersim.envs.intersimple import NRasterized
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
"""
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))
if __name__ == '__main__':
import fire
fire.Fire(render_env)