File size: 3,444 Bytes
bdce880 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | import torch
from tqdm import *
class IdentityTransformer():
def __init__(self, X):
self.mean = X.mean(dim=0, keepdim=True)
self.std = X.std(dim=0, keepdim=True) + 1e-8
def to(self, device):
self.mean = self.mean.to(device)
self.std = self.std.to(device)
return self
def cuda(self):
self.mean = self.mean.cuda()
self.std = self.std.cuda()
def cpu(self):
self.mean = self.mean.cpu()
self.std = self.std.cpu()
def encode(self, x):
return x
def decode(self, x):
return x
class UnitTransformer():
def __init__(self, X):
self.mean = X.mean(dim=(0, 1), keepdim=True)
self.std = X.std(dim=(0, 1), keepdim=True) + 1e-8
def to(self, device):
self.mean = self.mean.to(device)
self.std = self.std.to(device)
return self
def cuda(self):
self.mean = self.mean.cuda()
self.std = self.std.cuda()
def cpu(self):
self.mean = self.mean.cpu()
self.std = self.std.cpu()
def encode(self, x):
x = (x - self.mean) / (self.std)
return x
def decode(self, x):
return x * self.std + self.mean
def transform(self, X, inverse=True, component='all'):
if component == 'all' or 'all-reduce':
if inverse:
orig_shape = X.shape
return (X * (self.std - 1e-8) + self.mean).view(orig_shape)
else:
return (X - self.mean) / self.std
else:
if inverse:
orig_shape = X.shape
return (X * (self.std[:, component] - 1e-8) + self.mean[:, component]).view(orig_shape)
else:
return (X - self.mean[:, component]) / self.std[:, component]
class UnitGaussianNormalizer(object):
def __init__(self, x, eps=0.00001, time_last=True):
super(UnitGaussianNormalizer, self).__init__()
self.mean = torch.mean(x, 0)
self.std = torch.std(x, 0)
self.eps = eps
self.time_last = time_last # if the time dimension is the last dim
def encode(self, x):
x = (x - self.mean) / (self.std + self.eps)
return x
def decode(self, x, sample_idx=None):
# sample_idx is the spatial sampling mask
if sample_idx is None:
std = self.std + self.eps # n
mean = self.mean
else:
if self.mean.ndim == sample_idx.ndim or self.time_last:
std = self.std[sample_idx] + self.eps # batch*n
mean = self.mean[sample_idx]
if self.mean.ndim > sample_idx.ndim and not self.time_last:
std = self.std[..., sample_idx] + self.eps # T*batch*n
mean = self.mean[..., sample_idx]
# x is in shape of batch*(spatial discretization size) or T*batch*(spatial discretization size)
x = (x * std) + mean
return x
def to(self, device):
if torch.is_tensor(self.mean):
self.mean = self.mean.to(device)
self.std = self.std.to(device)
else:
self.mean = torch.from_numpy(self.mean).to(device)
self.std = torch.from_numpy(self.std).to(device)
return self
def cuda(self):
self.mean = self.mean.cuda()
self.std = self.std.cuda()
def cpu(self):
self.mean = self.mean.cpu()
self.std = self.std.cpu()
|