mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-10 11:40:58 +08:00
Support deterministic test episode
This commit is contained in:
@@ -56,3 +56,5 @@ class Config:
|
||||
self.gaussian_noise_scale = 0.3
|
||||
self.optimization_epochs = 4
|
||||
self.num_mini_batches = 32
|
||||
self.test_interval = 0
|
||||
self.test_repetitions = 10
|
||||
|
||||
+11
-1
@@ -16,6 +16,7 @@ def run_episodes(agent):
|
||||
ep = 0
|
||||
rewards = []
|
||||
steps = []
|
||||
avg_test_rewards = []
|
||||
agent_type = agent.__class__.__name__
|
||||
while True:
|
||||
ep += 1
|
||||
@@ -38,8 +39,17 @@ def run_episodes(agent):
|
||||
if config.max_steps and agent.total_steps > config.max_steps:
|
||||
break
|
||||
|
||||
if config.test_interval and ep % config.test_interval == 0:
|
||||
test_rewards = []
|
||||
for _ in range(config.test_repetitions):
|
||||
test_rewards.append(agent.episode(True)[0])
|
||||
avg_reward = np.mean(test_rewards)
|
||||
avg_test_rewards.append(avg_reward)
|
||||
config.logger.info('Averaged test reward %f(%f)' % (
|
||||
avg_reward, np.std(test_rewards) / np.sqrt(config.test_repetitions)))
|
||||
|
||||
agent.close()
|
||||
return steps, rewards
|
||||
return steps, rewards, avg_test_rewards
|
||||
|
||||
def run_iterations(agent):
|
||||
config = agent.config
|
||||
|
||||
Reference in New Issue
Block a user