Fix predict for reparameterized modules
Better way would probably be to rewrite flat_grad and reparam
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user