sayest-verifier / models.py
Aleksei Žavoronkov
update model architecture to the latest
8cf218e
Raw
History Blame Contribute Delete
2.68 kB
import torch
import torch.nn as nn
class OrdinalLogLoss(nn.Module):
def __init__(
self,
num_classes,
alpha=1.0,
reduction='mean',
distance_matrix=None,
class_weights=None,
eps=1e-8,
ignore_index=-100,
):
super(OrdinalLogLoss, self).__init__()
self.num_classes = num_classes
self.alpha = alpha
self.reduction = reduction
self.eps = eps
self.ignore_index = ignore_index
if distance_matrix is not None:
assert distance_matrix.shape == (num_classes, num_classes), \
"Distance matrix must be of shape (num_classes, num_classes)"
self.register_buffer('distance_matrix', distance_matrix.float())
else:
idx = torch.arange(num_classes).float()
default_matrix = torch.abs(idx.unsqueeze(0) - idx.unsqueeze(1))
self.register_buffer('distance_matrix', default_matrix)
if class_weights is not None:
assert class_weights.shape == (num_classes,), \
"Class weights must be of shape (num_classes,)"
self.register_buffer('class_weights', class_weights.float())
else:
self.class_weights = None
def forward(self, logits, target):
if logits.numel() == 0:
return logits.new_tensor(0.0)
probs = torch.softmax(logits, dim=-1).clamp(max=1 - self.eps)
if self.ignore_index is not None:
valid_mask = target != self.ignore_index
else:
valid_mask = torch.ones_like(target, dtype=torch.bool)
if not valid_mask.any():
if self.reduction == 'none':
return logits.new_zeros(target.shape, dtype=logits.dtype)
return logits.new_tensor(0.0)
active_probs = probs[valid_mask]
active_target = target[valid_mask]
distances = self.distance_matrix[active_target] ** self.alpha
per_class_loss = -torch.log(1 - active_probs + self.eps)
loss_active = (per_class_loss * distances).sum(dim=-1)
if self.class_weights is not None:
sample_weights = self.class_weights[active_target]
loss_active = loss_active * sample_weights
if self.reduction == 'none':
full_loss = logits.new_zeros(target.shape, dtype=logits.dtype)
full_loss[valid_mask] = loss_active
return full_loss
if self.reduction == 'mean':
return loss_active.mean()
elif self.reduction == 'sum':
return loss_active.sum()
else:
raise ValueError(f"Unsupported reduction: {self.reduction}")