mirror of
https://github.com/wassname/DeepRL.git
synced 2026-08-30 11:14:19 +08:00
Deterministic episodes
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user