mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Refactor DDPG
This commit is contained in:
+1
-2
@@ -36,8 +36,6 @@ class Config:
|
||||
self.gae_tau = 1.0
|
||||
self.noise_decay_interval = 0
|
||||
self.target_network_mix = 0.001
|
||||
self.reward_shift_fn = lambda r: r
|
||||
self.state_shift_fn = lambda s: s
|
||||
self.action_shift_fn = lambda a: a
|
||||
self.reward_weight = 1
|
||||
self.hybrid_reward = False
|
||||
@@ -47,3 +45,4 @@ class Config:
|
||||
self.master_fn = None
|
||||
self.master_optimizer_fn = None
|
||||
self.num_heads = 10
|
||||
self.min_epsilon = 0
|
||||
|
||||
@@ -5,6 +5,17 @@
|
||||
#######################################################################
|
||||
import torch
|
||||
|
||||
class Normalizer:
|
||||
def __init__(self, o_size):
|
||||
self.stats = SharedStats(o_size)
|
||||
|
||||
def __call__(self, o_):
|
||||
o = torch.FloatTensor(o_)
|
||||
self.stats.feed(o)
|
||||
std = (self.stats.v + 1e-6) ** .5
|
||||
o = (o - self.stats.m) / std
|
||||
return o.numpy().reshape(o_.shape)
|
||||
|
||||
class StaticNormalizer:
|
||||
def __init__(self, o_size):
|
||||
self.offline_stats = SharedStats(o_size)
|
||||
|
||||
Reference in New Issue
Block a user