Deterministic episodes

This commit is contained in:
Shangtong Zhang
2018-05-29 11:02:37 -06:00
parent 89b1450df7
commit 5a1d489938
3 changed files with 26 additions and 3 deletions
+21 -1
View File
@@ -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)
+3 -2
View File
@@ -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
+2
View File
@@ -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)