fix training

This commit is contained in:
Andras Schmelczer 2024-04-13 15:54:34 +01:00
commit 4349f5af41
No known key found for this signature in database
GPG key ID: FC8F2C3D3D1A718C
4 changed files with 85068 additions and 118911 deletions

View file

@ -1 +1,3 @@
from .histogram_dataset import HistogramDataset from .histogram_dataset import HistogramDataset
from .random_edit import random_edit
from .progressive_pooling_loss import ProgressivePoolingLoss

View file

@ -61,13 +61,14 @@ class HistogramDataset(Dataset):
edited = random_edit(original, seed=idx) edited = random_edit(original, seed=idx)
original_histogram = compute_histogram(
original, bins=self._bin_count, normalize=True
)
edited_histogram = compute_histogram( edited_histogram = compute_histogram(
edited, bins=self._bin_count, normalize=True edited, bins=self._bin_count, normalize=True
) )
original_histogram = compute_histogram(
original, bins=self._bin_count, normalize=True
)
result = ( result = (
torch.tensor(edited_histogram, dtype=torch.float).unsqueeze(0), torch.tensor(edited_histogram, dtype=torch.float).unsqueeze(0),
torch.tensor(original_histogram, dtype=torch.float).unsqueeze(0), torch.tensor(original_histogram, dtype=torch.float).unsqueeze(0),

View file

@ -1,31 +1,38 @@
from typing import List
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
class ProgressivePoolingLoss(nn.Module): class ProgressivePoolingLoss(nn.Module):
def __init__(self, initial_pool_size: int = 2, damping=1.8): def __init__(self, target_sizes: List[int], damping: float):
super(ProgressivePoolingLoss, self).__init__() super(ProgressivePoolingLoss, self).__init__()
self._initial_pool_size = initial_pool_size self._target_sizes = target_sizes
self._damping = damping self._damping = damping
def forward(self, tensor_a, tensor_b): def forward(self, tensor_a, tensor_b):
assert ( assert (
tensor_a.size() == tensor_b.size() tensor_a.size() == tensor_b.size()
), "Input tensors must have the same size." ), f"Input tensors must have the same size, got {tensor_a.size()} and {tensor_b.size()}"
max_pool_size = min(tensor_a.size(1), tensor_a.size(2), tensor_a.size(3)) assert (
len(tensor_a.size()) == 5
), f"Input tensors must have 5 dimensions, got {tensor_a.size()}"
_minibatch_size, _channels, depth, height, width = tensor_a.size()
assert depth == height == width, "Input tensors must be cubes."
loss = 0.0 loss = 0.0
damping = 1 weight = 1
for pool_size in range(self._initial_pool_size, max_pool_size): for target_size in self._target_sizes:
pool_size = depth // target_size
pooled_a = F.avg_pool3d(tensor_a, pool_size) * (pool_size**3) pooled_a = F.avg_pool3d(tensor_a, pool_size) * (pool_size**3)
pooled_b = F.avg_pool3d(tensor_b, pool_size) * (pool_size**3) pooled_b = F.avg_pool3d(tensor_b, pool_size) * (pool_size**3)
diff = torch.square(pooled_a - pooled_b) diff = torch.abs(pooled_a - pooled_b)
loss += diff.mean() / damping loss += diff.mean() * weight
damping *= self._damping weight *= self._damping
return loss return loss

206249
train.ipynb

File diff suppressed because one or more lines are too long