| 340 | |
| 341 | # Maybe not necessary in image inputs |
| 342 | class 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.) |
nothing calls this directly
no outgoing calls
no test coverage detected