# Focal loss weights for preffered labels FOCAL_ALPHA_DEFAULT = 0.25 FOCAL_ALPHA_PREFERRED = 0.75 PREFERRED_LABELS = { "fear", "sadness", "disgust", "disapproval", "annoyance", "anger", "disappointment", "optimism", "amusement", "surprise", "admiration", "excitement", "confusion","joy","love" } FOCAL_ALPHA_PER_LABEL: list[float] = [ FOCAL_ALPHA_PREFERRED if lbl in PREFERRED_LABELS else FOCAL_ALPHA_DEFAULT for lbl in EMOTION_LABELS ] "Per-label weighted focal binary cross-entropy for multi-label problems" class FocalLossWithAlpha(nn.Module): def __init__(self, alpha: list[float], gamma: float = 2.0): super().__init__() self.register_buffer("alpha", torch.tensor(alpha, dtype=torch.float32)) self.gamma = gamma def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: probs = torch.sigmoid(logits) p_t = probs * targets + (1.0 - probs) * (1.0 - targets) alpha_t = self.alpha * targets + (1.0 - self.alpha) * (1.0 - targets) focal_w = alpha_t * (1.0 - p_t) ** self.gamma bce = nn.functional.binary_cross_entropy_with_logits( logits, targets, reduction="none" ) return (focal_w * bce).mean()