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

@@ -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':