mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Continuous A3C
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
from config import *
|
||||
from shifter import *
|
||||
try:
|
||||
from tf_logger import Logger
|
||||
except:
|
||||
|
||||
@@ -8,6 +8,7 @@ class Config:
|
||||
def __init__(self):
|
||||
self.task_fn = None
|
||||
self.optimizer_fn = None
|
||||
self.critic_optimizer_fn = None
|
||||
self.network_fn = None
|
||||
self.policy_fn = None
|
||||
self.replay_fn = None
|
||||
@@ -27,3 +28,6 @@ class Config:
|
||||
self.gradient_clip = 40
|
||||
self.entropy_weight = 0.01
|
||||
self.gae_tau = 1.0
|
||||
self.reward_shift_fn = lambda r: r
|
||||
self.state_shift_fn = lambda s: s
|
||||
self.action_shift_fn = lambda a: a
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
# Adapted from https://github.com/kvfrans/parallel-trpo/blob/master/utils.py
|
||||
class Shifter:
|
||||
def __init__(self, filter_mean=True):
|
||||
self.m = 0
|
||||
self.v = 0
|
||||
self.n = 0.
|
||||
self.filter_mean = filter_mean
|
||||
|
||||
def state_dict(self):
|
||||
return {'m': self.m,
|
||||
'v': self.v,
|
||||
'n': self.n}
|
||||
|
||||
def load_state_dict(self, saved):
|
||||
self.m = saved['m']
|
||||
self.v = saved['v']
|
||||
self.n = saved['n']
|
||||
|
||||
def __call__(self, o):
|
||||
self.m = self.m * (self.n / (self.n + 1)) + o * 1 / (1 + self.n)
|
||||
self.v = self.v * (self.n / (self.n + 1)) + (o - self.m) ** 2 * 1 / (1 + self.n)
|
||||
self.std = (self.v + 1e-6) ** .5 # std
|
||||
self.n += 1
|
||||
if self.filter_mean:
|
||||
o_ = (o - self.m) / self.std
|
||||
else:
|
||||
o_ = o / self.std
|
||||
return o_
|
||||
Reference in New Issue
Block a user