Fix predict for reparameterized modules

Better way would probably be to rewrite flat_grad and reparam
This commit is contained in:
ebuehrle
2022-02-21 13:28:25 +01:00
parent 2da0e05782
commit a242edc5d3
4 changed files with 25 additions and 12 deletions

View File

@@ -17,10 +17,13 @@ class BasePolicy(nn.Module):
def predict(self, observations, state=None, episode_start=None, deterministic=True): def predict(self, observations, state=None, episode_start=None, deterministic=True):
observations = torch.tensor(observations) observations = torch.tensor(observations)
return self._predict(self.forward(observations), state, episode_start, deterministic)
def _predict(self, dist, state=None, episode_start=None, deterministic=True):
if deterministic: if deterministic:
actions = self.forward(observations)[..., :self.action_dim] actions = dist[..., :self.action_dim]
else: else:
actions = self.sample(self.forward(observations)) actions = self.sample(dist)
return actions, None return actions, None
def log_prob(self, dist, actions): def log_prob(self, dist, actions):
@@ -64,12 +67,11 @@ class DiscretePolicy(BasePolicy):
def torch_dist(self, dist): def torch_dist(self, dist):
return Categorical(logits=dist) return Categorical(logits=dist)
def predict(self, observations, state=None, episode_start=None, deterministic=True): def _predict(self, dist, state=None, episode_start=None, deterministic=True):
observations = torch.tensor(observations)
if deterministic: if deterministic:
_, actions = self.forward(observations).max(-1) _, actions = dist.max(-1)
else: else:
actions = self.sample(self.forward(observations)) actions = self.sample(dist)
return actions, None return actions, None
class SetPolicy(Policy): class SetPolicy(Policy):

View File

@@ -158,8 +158,16 @@ class ReparamPolicy(ReparamModule):
def kl_divergence(self, *args, **kwargs): def kl_divergence(self, *args, **kwargs):
return self.module.kl_divergence(*args, **kwargs) return self.module.kl_divergence(*args, **kwargs)
def predict(self, *args, **kwargs): def predict(self, obs, *args, **kwargs):
return self.module.predict(*args, **kwargs) obs = torch.tensor(obs)
return self.module._predict(self.forward(obs), *args, **kwargs)
def unsafe_probability_mass(self, *args, **kwargs): def unsafe_probability_mass(self, *args, **kwargs):
return self.module.unsafe_probability_mass(*args, **kwargs) return self.module.unsafe_probability_mass(*args, **kwargs)
class ReparamSafePolicy(ReparamPolicy):
def predict(self, obs, *args, **kwargs):
observation = torch.tensor(obs['observation'])
safe_actions = torch.tensor(obs['safe_actions'])
return self.module._predict(self.forward(observation, safe_actions), *args, **kwargs)

View File

@@ -10,7 +10,7 @@ import src.gail.options as options_envs
from src.evaluation.metrics import divergence, visualize_distribution, rwse from src.evaluation.metrics import divergence, visualize_distribution, rwse
from src.evaluation.utils import save_metrics from src.evaluation.utils import save_metrics
from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.policy import SetPolicy, SetDiscretePolicy
from src.core.reparam_module import ReparamPolicy from src.core.reparam_module import ReparamPolicy, ReparamSafePolicy
from src.options import envs as options_envs2 from src.options import envs as options_envs2
from src.safe_options.policy import SetMaskedDiscretePolicy from src.safe_options.policy import SetMaskedDiscretePolicy
from src.safe_options import options as options_envs3 from src.safe_options import options as options_envs3
@@ -70,7 +70,7 @@ def load_policy(method:str,
torch.zeros(env.observation_space['observation'].shape), torch.zeros(env.observation_space['observation'].shape),
torch.zeros(env.observation_space['safe_actions'].shape) torch.zeros(env.observation_space['safe_actions'].shape)
) )
policy = ReparamPolicy(policy) policy = ReparamSafePolicy(policy)
policy.load_state_dict(torch.load(policy_file)) policy.load_state_dict(torch.load(policy_file))
policy.eval() policy.eval()
elif method == 'sgail-ppo': elif method == 'sgail-ppo':

View File

@@ -24,10 +24,13 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy):
def predict(self, observations, state=None, episode_start=None, deterministic=True): def predict(self, observations, state=None, episode_start=None, deterministic=True):
observation = torch.tensor(observations['observation']) observation = torch.tensor(observations['observation'])
safe_actions = torch.tensor(observations['safe_actions']) safe_actions = torch.tensor(observations['safe_actions'])
return self._predict(self.forward(observation, safe_actions), state, episode_start, deterministic)
def _predict(self, dist, state=None, episode_start=None, deterministic=True):
if deterministic: if deterministic:
_, actions = self.forward(observation, safe_actions).max(-1) _, actions = self.torch_dist(dist).probs.max(-1)
else: else:
actions = self.sample(self.forward(observation, safe_actions)) actions = self.sample(dist)
return actions, None return actions, None
# def torch_dist_nomask(self, dist): # def torch_dist_nomask(self, dist):