load_state_dict for Normalizer, StaticNormalizer

This commit is contained in:
Mike Clark
2018-02-02 15:14:49 +08:00
committed by GitHub
parent 63c5d5ea54
commit db94825b19
+13 -1
View File
@@ -24,6 +24,12 @@ class Normalizer:
else:
o = o.reshape(o_.shape)
return o
def state_dict(self):
return self.stats.state_dict()
def load_state_dict(self, saved):
self.stats.load_state_dict(saved)
class StaticNormalizer:
def __init__(self, o_size):
@@ -46,6 +52,12 @@ class StaticNormalizer:
else:
o = o.reshape(o_.shape)
return o
def state_dict(self):
return self.offline_stats.state_dict()
def load_state_dict(self, saved):
self.offline_stats.load_state_dict(saved)
class SharedStats:
def __init__(self, o_size):
@@ -94,4 +106,4 @@ class SharedStats:
def load_state_dict(self, saved):
self.m = torch.FloatTensor(saved['m'])
self.v = torch.FloatTensor(saved['v'])
self.n = torch.FloatTensor(saved['n'])
self.n = torch.FloatTensor(saved['n'])