diff --git a/scripts/algorithms/td3/agent.py b/scripts/algorithms/td3/agent.py index 882e1e4..fadf64f 100644 --- a/scripts/algorithms/td3/agent.py +++ b/scripts/algorithms/td3/agent.py @@ -65,9 +65,10 @@ class Agent(AbstractAgent): """ AbstractAgent.__init__(self, env, args) - self.actor, self.actor_target, self.critic1, self.critic1_target, self.critic2, self.critic2_target = ( # noqa: B950 - models - ) + self.actor, self.actor_target = models[:2] + self.critic1, self.critic1_target = models[2:4] + self.critic2, self.critic2_target = models[4:] + self.actor_optim, self.critic_optim = optims self.hyper_params = hyper_params self.exploration_noise, self.target_policy_noise = noises diff --git a/scripts/data/reacher_demo.pkl b/scripts/data/reacher_demo.pkl new file mode 100644 index 0000000..0cd3e88 Binary files /dev/null and b/scripts/data/reacher_demo.pkl differ diff --git a/scripts/run_lunarlander_continuous.py b/scripts/run_lunarlander_continuous.py index 55987a2..3a6766c 100644 --- a/scripts/run_lunarlander_continuous.py +++ b/scripts/run_lunarlander_continuous.py @@ -37,6 +37,13 @@ parser.add_argument( default=0, help="start rendering after the input number of episode", ) + +parser.add_argument( + "--demo-path", + type=str, + default="data/lunarlander_continuous_demo.pkl", + help="demonstration path", +) parser.add_argument("--save-period", type=int, default=100, help="save model period") parser.add_argument("--log", action="store_true", help="turn on logging") parser.add_argument("--test", action="store_true", help="test mode (no training)") diff --git a/scripts/run_reacher_v2.py b/scripts/run_reacher_v2.py index c4ea8eb..5b3f37a 100644 --- a/scripts/run_reacher_v2.py +++ b/scripts/run_reacher_v2.py @@ -40,10 +40,7 @@ parser.add_argument( "--max-episode-steps", type=int, default=-1, help="max episode step" ) parser.add_argument( - "--demo-path", - type=str, - default="data/lunarlander_continuous_demo.pkl", - help="demonstration path", + "--demo-path", type=str, default="data/reacher_demo.pkl", help="demonstration path" ) parser.set_defaults(test=False)