[文档]classValueNorm(nn.Module):""" Normalize a vector of observations - across the first norm_axes dimensions"""def__init__(self,input_shape,norm_axes=1,beta=0.99999,per_element_update=False,epsilon=1e-5):super(ValueNorm,self).__init__()self.input_shape=input_shapeself.norm_axes=norm_axesself.epsilon=epsilonself.beta=betaself.per_element_update=per_element_updateself.running_mean=nn.Parameter(torch.zeros(input_shape),requires_grad=False)self.running_mean_sq=nn.Parameter(torch.zeros(input_shape),requires_grad=False)self.debiasing_term=nn.Parameter(torch.tensor(0.0),requires_grad=False)self.reset_parameters()
@torch.no_grad()defupdate(self,input_vector):iftype(input_vector)==np.ndarray:input_vector=torch.from_numpy(input_vector)input_vector=input_vector.to(self.running_mean.device)# not elegant, but works in most casesbatch_mean=input_vector.mean(dim=tuple(range(self.norm_axes)))batch_sq_mean=(input_vector**2).mean(dim=tuple(range(self.norm_axes)))ifself.per_element_update:batch_size=np.prod(input_vector.size()[:self.norm_axes])weight=self.beta**batch_sizeelse:weight=self.betaself.running_mean.mul_(weight).add_(batch_mean*(1.0-weight))self.running_mean_sq.mul_(weight).add_(batch_sq_mean*(1.0-weight))self.debiasing_term.mul_(weight).add_(1.0*(1.0-weight))
[文档]defnormalize(self,input_vector):# Make sure input is float32iftype(input_vector)==np.ndarray:input_vector=torch.from_numpy(input_vector)input_vector=input_vector.to(self.running_mean.device)# not elegant, but works in most casesmean,var=self.running_mean_var()out=(input_vector-mean[(None,)*self.norm_axes])/torch.sqrt(var)[(None,)*self.norm_axes]returnout
[文档]defdenormalize(self,input_vector):""" Transform normalized data back into original distribution """input_vector=torch.as_tensor(input_vector)input_vector=input_vector.to(self.running_mean.device)# not elegant, but works in most casesmean,var=self.running_mean_var()out=input_vector*torch.sqrt(var)[(None,)*self.norm_axes]+mean[(None,)*self.norm_axes]out=out.cpu().numpy()returnout