diff --git a/requirements.txt b/requirements.txt index 62870bb..43f946b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,4 +7,5 @@ tqdm ray[tune] hyperopt psutil -fire \ No newline at end of file +fire +stable_baselines3 diff --git a/src/eval_main.py b/src/eval_main.py index 7fc5885..3f59844 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -37,21 +37,22 @@ def load_policy(method:str, Returns: policy (Optional[BaseAlgorithm]): the policy to evaluate """ + ml = torch.device('cpu') if not torch.cuda.is_available() else None if method == 'idm': policy = IDMRulePolicy(env, **policy_kwargs) elif method == 'bc': policy = SetPolicy(env.action_space.shape[-1]) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'gail': policy = SetPolicy(env.action_space.shape[-1]) policy(torch.zeros(env.observation_space.shape)) policy = ReparamPolicy(policy) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'gail-ppo': policy = SetPolicy(env.action_space.shape[-1]) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'rail': raise NotImplementedError @@ -59,11 +60,11 @@ def load_policy(method:str, policy = SetDiscretePolicy(env.action_space.n) policy(torch.zeros(env.observation_space.shape)) policy = ReparamPolicy(policy) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'ogail-ppo': policy = SetDiscretePolicy(env.action_space.n) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'sgail': policy = SetMaskedDiscretePolicy(env.action_space.n) @@ -72,11 +73,11 @@ def load_policy(method:str, torch.zeros(env.observation_space['safe_actions'].shape) ) policy = ReparamSafePolicy(policy) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() elif method == 'sgail-ppo': policy = SetMaskedDiscretePolicy(env.action_space.n) - policy.load_state_dict(torch.load(policy_file)) + policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.eval() else: raise NotImplementedError @@ -406,4 +407,4 @@ def eval_main( if __name__=='__main__': import fire - fire.Fire(eval_main) \ No newline at end of file + fire.Fire(eval_main)