Spaces:
Running on Zero
Running on Zero
File size: 8,629 Bytes
d1f8fab | 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 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | from __future__ import annotations
import math
from typing import Dict
import torch
import torch.nn as nn
import torch.nn.functional as F
def _off_diagonal_mask(size: int, device: torch.device) -> torch.Tensor:
mask = torch.ones(size, size, dtype=torch.bool, device=device)
mask.fill_diagonal_(False)
return mask
def supervised_contrastive_loss(embeddings: torch.Tensor, labels: torch.Tensor, temperature: float) -> torch.Tensor:
if embeddings.size(0) <= 1:
return embeddings.new_zeros(())
normalized = F.normalize(embeddings.float(), dim=-1)
logits = torch.matmul(normalized, normalized.t()) / temperature
logits_mask = ~torch.eye(logits.size(0), dtype=torch.bool, device=logits.device)
logits = logits - logits.max(dim=1, keepdim=True).values.detach()
label_mask = labels.unsqueeze(0).eq(labels.unsqueeze(1)) & logits_mask
positives_per_row = label_mask.sum(dim=1)
valid_rows = positives_per_row > 0
if not valid_rows.any():
return embeddings.new_zeros(())
exp_logits = torch.exp(logits) * logits_mask.to(dtype=logits.dtype)
log_prob = logits - torch.log(exp_logits.sum(dim=1, keepdim=True).clamp_min(1e-12))
positive_log_prob = log_prob.masked_fill(~label_mask, 0.0)
loss = -positive_log_prob.sum(dim=1) / positives_per_row.clamp_min(1)
return loss[valid_rows].mean().to(dtype=embeddings.dtype)
class MultiPrototypeLoss(nn.Module):
def __init__(
self,
classification_weight: float = 1.0,
arcface_scale: float = 30.0,
arcface_margin: float = 0.30,
compactness_weight: float = 0.25,
diversity_weight: float = 0.05,
usage_weight: float = 0.05,
branch_orthogonality_weight: float = 0.05,
branch_balance_weight: float = 0.05,
view_balance_weight: float = 0.05,
supcon_weight: float = 0.1,
label_smoothing: float = 0.0,
assignment_temperature: float = 0.1,
supcon_temperature: float = 0.1,
prototype_margin: float = 0.15,
) -> None:
super().__init__()
self.classification_weight = classification_weight
self.arcface_scale = float(arcface_scale)
self.arcface_margin = float(arcface_margin)
self.compactness_weight = compactness_weight
self.diversity_weight = diversity_weight
self.usage_weight = usage_weight
self.branch_orthogonality_weight = branch_orthogonality_weight
self.branch_balance_weight = branch_balance_weight
self.view_balance_weight = view_balance_weight
self.supcon_weight = supcon_weight
self.label_smoothing = label_smoothing
self.assignment_temperature = assignment_temperature
self.supcon_temperature = supcon_temperature
self.prototype_margin = prototype_margin
def _arcface_loss(self, cosine: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
cosine = cosine.float().clamp(-1.0 + 1e-6, 1.0 - 1e-6)
sine = torch.sqrt((1.0 - cosine.pow(2)).clamp_min(1e-6))
cos_m = math.cos(self.arcface_margin)
sin_m = math.sin(self.arcface_margin)
th = math.cos(math.pi - self.arcface_margin)
mm = math.sin(math.pi - self.arcface_margin) * self.arcface_margin
phi = cosine * cos_m - sine * sin_m
phi = torch.where(cosine > th, phi, cosine - mm)
one_hot = torch.zeros_like(cosine)
one_hot.scatter_(1, labels.view(-1, 1), 1.0)
logits = (one_hot * phi + (1.0 - one_hot) * cosine) * self.arcface_scale
return F.cross_entropy(logits, labels, label_smoothing=self.label_smoothing)
def _compactness_loss(self, true_sims: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
assignments = torch.softmax(true_sims / self.assignment_temperature, dim=-1)
compactness = (1.0 - (assignments * true_sims).sum(dim=-1)).mean()
return compactness, assignments
def _prototype_diversity_loss(self, normalized_prototypes: torch.Tensor) -> torch.Tensor:
pairwise = torch.matmul(normalized_prototypes, normalized_prototypes.transpose(-1, -2))
mask = _off_diagonal_mask(pairwise.size(-1), pairwise.device)
off_diag = pairwise[:, mask].view(pairwise.size(0), -1)
return F.relu(off_diag - self.prototype_margin).pow(2).mean()
def _prototype_usage_loss(self, assignments: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
losses = []
target = torch.full(
(assignments.size(-1),),
1.0 / assignments.size(-1),
device=assignments.device,
dtype=assignments.dtype,
)
for label in labels.unique():
mask = labels == label
if mask.sum() < 2:
continue
mean_assignment = assignments[mask].mean(dim=0)
losses.append((mean_assignment - target).pow(2).mean())
if not losses:
return assignments.new_zeros(())
return torch.stack(losses).mean()
def _branch_orthogonality_loss(self, branch_embeddings: torch.Tensor) -> torch.Tensor:
normalized = F.normalize(branch_embeddings, dim=-1)
gram = torch.matmul(normalized, normalized.transpose(-1, -2))
mask = _off_diagonal_mask(gram.size(-1), gram.device)
off_diag = gram[:, mask].view(gram.size(0), -1)
return off_diag.pow(2).mean()
def _distribution_balance_loss(self, weights: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
flat_weights = weights.reshape(-1, weights.size(-1))
flat_mask = mask.reshape(-1, mask.size(-1))
valid = flat_mask.sum(dim=-1) > 1
if not valid.any():
return weights.new_zeros(())
valid_weights = flat_weights[valid]
valid_mask = flat_mask[valid]
target = valid_mask / valid_mask.sum(dim=-1, keepdim=True).clamp_min(1.0)
sample_loss = ((valid_weights - target).pow(2) * valid_mask).sum(dim=-1)
sample_loss = sample_loss / valid_mask.sum(dim=-1).clamp_min(1.0)
batch_mean = valid_weights.mean(dim=0)
batch_target = target.mean(dim=0)
batch_loss = (batch_mean - batch_target).pow(2).mean()
return sample_loss.mean() + batch_loss
def forward(self, outputs: Dict[str, torch.Tensor], labels: torch.Tensor) -> Dict[str, torch.Tensor]:
logits = outputs["class_logits"]
arcface_cosine = outputs["arcface_cosine"]
prototype_sims = outputs["prototype_similarities"]
branch_embeddings = outputs["branch_embeddings"]
normalized_prototypes = outputs["normalized_prototypes"]
embeddings = outputs["embedding"]
branch_weights = outputs["branch_weights"]
effective_branch_mask = outputs["effective_branch_mask"]
stacked_view_weights = outputs["stacked_view_weights"]
effective_view_mask = outputs["effective_view_mask"]
classification = self._arcface_loss(arcface_cosine, labels)
true_sims = prototype_sims[torch.arange(labels.size(0), device=labels.device), labels]
compactness, assignments = self._compactness_loss(true_sims)
diversity = self._prototype_diversity_loss(normalized_prototypes)
usage = self._prototype_usage_loss(assignments, labels)
branch_orthogonality = self._branch_orthogonality_loss(branch_embeddings)
branch_balance = self._distribution_balance_loss(branch_weights, effective_branch_mask)
expanded_view_mask = effective_view_mask.unsqueeze(1).expand_as(stacked_view_weights)
view_balance = self._distribution_balance_loss(stacked_view_weights, expanded_view_mask)
supcon = supervised_contrastive_loss(embeddings, labels, temperature=self.supcon_temperature)
total = (
self.classification_weight * classification
+ self.compactness_weight * compactness
+ self.diversity_weight * diversity
+ self.usage_weight * usage
+ self.branch_orthogonality_weight * branch_orthogonality
+ self.branch_balance_weight * branch_balance
+ self.view_balance_weight * view_balance
+ self.supcon_weight * supcon
)
return {
"loss": total,
"classification": classification.detach(),
"compactness": compactness.detach(),
"diversity": diversity.detach(),
"usage": usage.detach(),
"branch_orthogonality": branch_orthogonality.detach(),
"branch_balance": branch_balance.detach(),
"view_balance": view_balance.detach(),
"supcon": supcon.detach(),
}
|