GPUCodeForces/S1/uucoco_#90/HistogramLoss_torch.py

51 lines
1.5 KiB
Python

import torch
import torch.nn as nn
class Model(nn.Module):
def __init__(self, num_bins=10, min_val=0.0, max_val=1.0):
super().__init__()
self.num_bins = num_bins
self.min_val = min_val
self.max_val = max_val
self.step = (max_val - min_val) / num_bins
self.centers = torch.linspace(min_val + self.step / 2, max_val - self.step / 2, num_bins)
def forward(self, pos: torch.Tensor, neg: torch.Tensor) -> torch.Tensor:
delta = self.step
centers = self.centers.to(pos.device)
pos_rep = pos.unsqueeze(1).repeat(1, self.num_bins)
neg_rep = neg.unsqueeze(1).repeat(1, self.num_bins)
centers_rep_pos = centers.unsqueeze(0).repeat(pos.size(0), 1)
centers_rep_neg = centers.unsqueeze(0).repeat(neg.size(0), 1)
pos_hist = torch.clamp(1 - torch.abs(pos_rep - centers_rep_pos) / delta, min=0)
neg_hist = torch.clamp(1 - torch.abs(neg_rep - centers_rep_neg) / delta, min=0)
pos_hist_sum = pos_hist.sum(dim=0)
neg_hist_sum = neg_hist.sum(dim=0)
pos_cdf = torch.cumsum(pos_hist_sum, dim=0)
pos_cdf = pos_cdf / (pos_cdf[-1] + 1e-8)
neg_pdf = neg_hist_sum / (neg_hist_sum.sum() + 1e-8)
loss = (neg_pdf * pos_cdf).sum()
return loss
batch_size = 128
num_features = 512
def get_inputs():
pos = torch.rand(batch_size, dtype=torch.float32)
neg = torch.rand(batch_size, dtype=torch.float32)
return [pos, neg]
def get_init_inputs():
return [10, 0.0, 1.0]