This commit is contained in:
@@ -285,9 +285,9 @@ class Normalizer(nn.Module):
|
|||||||
new_xs = []
|
new_xs = []
|
||||||
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
|
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
|
||||||
for x, dataset_name in zip(xs, dataset_names):
|
for x, dataset_name in zip(xs, dataset_names):
|
||||||
# if dataset_name == "ex_normal":
|
delta = self.delta[dataset_name]
|
||||||
# dataset_name = "x2_normal"
|
delta = torch.where(delta == 0, torch.ones_like(delta), delta)
|
||||||
x = (x - self.min[dataset_name]) / (self.delta[dataset_name])
|
x = (x - self.min[dataset_name]) / delta
|
||||||
x = x * 2 - 1
|
x = x * 2 - 1
|
||||||
x = torch.clamp(x, -1, 1)
|
x = torch.clamp(x, -1, 1)
|
||||||
new_xs.append(x)
|
new_xs.append(x)
|
||||||
|
|||||||
Reference in New Issue
Block a user