File size: 1,807 Bytes
d4cbafd | 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 | import torch
### Helper functions for linear normalization and unnormalization
def normalize_to_neg_one_to_one(img):
return img * 2 - 1
def unnormalize_to_zero_to_one(t):
return (t + 1) * 0.5
def normalize_min_max(t, min_val, max_val, a, b, identity=False):
'''
Normalize t to [a, b] range
Args:
t: input tensor
min_val: minimum value of t
max_val: maximum value of t
a: minimum value of the output range
b: maximum value of the output range
'''
if identity:
return t
else:
return (b - a) * (t - min_val)/(max_val - min_val) + a
def unnormalize_min_max(t, min_val, max_val, a, b, identity=False):
'''
Unnormalize t from [a, b] range back to [min_val, max_val] range
Args:
t: input tensor
min_val: minimum value of t
max_val: maximum value of t
a: minimum value of the input range
b: maximum value of the input range
'''
if identity:
return t
else:
return (t - a) * (max_val - min_val)/(b - a) + min_val
def normalize_sqrt(traj_data, a, b):
'''
Normalize input tensor to [-1, 1] using square root.
@param traj_data: [*, 2]
@param a: [2]
@param b: [2]
'''
traj_data = torch.abs(traj_data).sqrt() * torch.sign(traj_data)
traj_data = traj_data / a.reshape(*([1] * (traj_data.dim() - 1)), -1) + b.reshape(*([1] * (traj_data.dim() - 1)), -1)
return traj_data
def unnormalize_sqrt(traj_data, a, b):
'''
Unnormalize input tensor from [-1, 1] using square root.
@param traj_data: [*, 2]
@param a: [2]
@param b: [2]
'''
traj_data = (traj_data - b.reshape(*([1] * (traj_data.dim() - 1)), -1)) * a.reshape(*([1] * (traj_data.dim() - 1)), -1)
traj_data = torch.sign(traj_data) * traj_data ** 2
return traj_data
|