diff --git a/src/core/policy.py b/src/core/policy.py index 5e377c9..5485159 100644 --- a/src/core/policy.py +++ b/src/core/policy.py @@ -17,10 +17,13 @@ class BasePolicy(nn.Module): def predict(self, observations, state=None, episode_start=None, deterministic=True): 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: - actions = self.forward(observations)[..., :self.action_dim] + actions = dist[..., :self.action_dim] else: - actions = self.sample(self.forward(observations)) + actions = self.sample(dist) return actions, None def log_prob(self, dist, actions): @@ -64,12 +67,11 @@ class DiscretePolicy(BasePolicy): def torch_dist(self, dist): return Categorical(logits=dist) - def predict(self, observations, state=None, episode_start=None, deterministic=True): - observations = torch.tensor(observations) + def _predict(self, dist, state=None, episode_start=None, deterministic=True): if deterministic: - _, actions = self.forward(observations).max(-1) + _, actions = dist.max(-1) else: - actions = self.sample(self.forward(observations)) + actions = self.sample(dist) return actions, None class SetPolicy(Policy): diff --git a/src/core/reparam_module.py b/src/core/reparam_module.py index 1b24986..a28f55e 100644 --- a/src/core/reparam_module.py +++ b/src/core/reparam_module.py @@ -158,8 +158,16 @@ class ReparamPolicy(ReparamModule): def kl_divergence(self, *args, **kwargs): return self.module.kl_divergence(*args, **kwargs) - def predict(self, *args, **kwargs): - return self.module.predict(*args, **kwargs) + def predict(self, obs, *args, **kwargs): + obs = torch.tensor(obs) + return self.module._predict(self.forward(obs), *args, **kwargs) def unsafe_probability_mass(self, *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) diff --git a/src/eval_main.py b/src/eval_main.py index c303f29..2e38932 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -10,7 +10,7 @@ import src.gail.options as options_envs from src.evaluation.metrics import divergence, visualize_distribution, rwse from src.evaluation.utils import save_metrics 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.safe_options.policy import SetMaskedDiscretePolicy 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['safe_actions'].shape) ) - policy = ReparamPolicy(policy) + policy = ReparamSafePolicy(policy) policy.load_state_dict(torch.load(policy_file)) policy.eval() elif method == 'sgail-ppo': diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py index 0325f59..9353a40 100644 --- a/src/safe_options/policy.py +++ b/src/safe_options/policy.py @@ -24,10 +24,13 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy): def predict(self, observations, state=None, episode_start=None, deterministic=True): observation = torch.tensor(observations['observation']) 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: - _, actions = self.forward(observation, safe_actions).max(-1) + _, actions = self.torch_dist(dist).probs.max(-1) else: - actions = self.sample(self.forward(observation, safe_actions)) + actions = self.sample(dist) return actions, None # def torch_dist_nomask(self, dist):