MCPcopy Create free account
hub / github.com/BIT-MCS/DRL-eFresh / Shared_obs_stats

Class Shared_obs_stats

methods/model.py:342–361  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

340
341# Maybe not necessary in image inputs
342class Shared_obs_stats():
343 def __init__(self, num_inputs,device):
344 self.n = torch.zeros(num_inputs).share_memory_().to(device)
345 self.mean = torch.zeros(num_inputs).share_memory_().to(device)
346 self.mean_diff = torch.zeros(num_inputs).share_memory_().to(device)
347 self.var = torch.zeros(num_inputs).share_memory_().to(device)
348
349 def observes(self, obs):
350 # observation mean var updates
351 x = obs.data.squeeze()
352 self.n += 1.
353 last_mean = self.mean.clone()
354 self.mean += (x - self.mean) / self.n
355 self.mean_diff += (x - last_mean) * (x - self.mean)
356 self.var = torch.clamp(self.mean_diff / self.n, min=1e-2)
357
358 def normalize(self, inputs):
359 obs_mean = self.mean.unsqueeze(0).expand_as(inputs)
360 obs_std = torch.sqrt(self.var).unsqueeze(0).expand_as(inputs)
361 return torch.clamp((inputs - obs_mean) / obs_std, -5., 5.)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected