diff --git a/deep_rl/agent/BaseAgent.py b/deep_rl/agent/BaseAgent.py index d5e8b07..c5a607d 100644 --- a/deep_rl/agent/BaseAgent.py +++ b/deep_rl/agent/BaseAgent.py @@ -6,6 +6,7 @@ import torch import numpy as np +from ..utils import * class BaseAgent: def __init__(self, config): @@ -35,9 +36,28 @@ class BaseAgent: self.config.state_normalizer.unset_read_only() return np.argmax(action.flatten()) + def deterministic_episode(self): + env = self.config.evaluation_env + state = env.reset() + total_rewards = 0 + while True: + action = self.evaluation_action(state) + state, reward, done, _ = env.step(action) + if done: + break + total_rewards += reward + return + + def evaluation_episodes(self): + interval = self.config.evaluation_episodes_interval + if not interval or self.total_steps % interval: + return + for ep in range(self.config.evaluation_episodes): + self.deterministic_episode() + def evaluate(self, steps=1): config = self.config - if config.evaluation_env is None: + if config.evaluation_env is None or self.config.evaluation_episodes_interval: return for _ in range(steps): action = self.evaluation_action(self.evaluation_state) diff --git a/deep_rl/agent/DDPG_agent.py b/deep_rl/agent/DDPG_agent.py index ead45ac..8269a46 100644 --- a/deep_rl/agent/DDPG_agent.py +++ b/deep_rl/agent/DDPG_agent.py @@ -44,6 +44,9 @@ class DDPGAgent(BaseAgent): steps = 0 total_reward = 0.0 while True: + self.evaluate() + self.evaluation_episodes() + action = self.network.predict(np.stack([state]), True).flatten() if not deterministic: action += self.random_process.sample() @@ -60,8 +63,6 @@ class DDPGAgent(BaseAgent): steps += 1 state = next_state - self.evaluate() - if not deterministic and self.replay.size() >= config.min_memory_size: experiences = self.replay.sample() states, actions, rewards, next_states, terminals = experiences diff --git a/deep_rl/utils/config.py b/deep_rl/utils/config.py index 590a576..b8cf743 100644 --- a/deep_rl/utils/config.py +++ b/deep_rl/utils/config.py @@ -61,6 +61,8 @@ class Config: self.test_repetitions = 10 self.evaluation_env = None self.termination_regularizer = 0 + self.evaluation_episodes_interval = 0 + self.evaluation_episodes = 0 def add_argument(self, *args, **kwargs): self.parser.add_argument(*args, **kwargs)