Files
modeling_zhaocui/deep_learning/Q1/q1/losses.py
T

176 lines
6.7 KiB
Python

from __future__ import annotations
from collections.abc import Mapping
import torch
import torch.nn.functional as F
from torch import Tensor
from .metrics import attention_row_similarity
from .types import AlignmentOutput, MODALITIES
def temporal_monotonicity_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
epsilon: float = 0.02,
) -> Tensor:
"""Penalize backward motion on the normalized clip timeline."""
if epsilon < 0:
raise ValueError("epsilon must be non-negative")
losses = []
for name in MODALITIES:
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
backward = F.relu(mu[:, :-1] - mu[:, 1:] - epsilon)
losses.append(backward.square().mean())
return torch.stack(losses).mean()
def cross_modal_contrastive_loss(
aligned: Mapping[str, Tensor], temperature: float = 0.1
) -> Tensor:
"""Symmetric in-batch InfoNCE over same-sample, same-grid-slot positives."""
if temperature <= 0:
raise ValueError("temperature must be positive")
if set(aligned) != set(MODALITIES):
raise ValueError(f"aligned must contain exactly {MODALITIES}")
pair_losses = []
for left_index, left_name in enumerate(MODALITIES):
for right_name in MODALITIES[left_index + 1 :]:
left = F.normalize(aligned[left_name].flatten(0, 1), dim=-1)
right = F.normalize(aligned[right_name].flatten(0, 1), dim=-1)
if left.shape != right.shape:
raise ValueError("contrastive representations must share [B, K, D]")
logits = left @ right.T / temperature
labels = torch.arange(logits.shape[0], device=logits.device)
pair_losses.append(
(F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2
)
return torch.stack(pair_losses).mean()
def masked_reconstruction_loss(prediction: Tensor, target: Tensor, mask: Tensor) -> Tensor:
"""Smooth-L1 loss over masked grid slots only."""
if prediction.shape != target.shape:
raise ValueError("prediction and target must have the same shape")
if mask.shape != target.shape[:2] or mask.dtype != torch.bool:
raise ValueError("mask must be boolean with shape [B, K]")
if not bool(mask.any()):
raise ValueError("mask must select at least one target slot")
element_loss = F.smooth_l1_loss(prediction, target, reduction="none")
return element_loss[mask].mean()
def temporal_span_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
minimum_span: float = 0.7,
modalities: tuple[str, ...] = MODALITIES,
) -> Tensor:
"""Penalize an expected-time path that does not cover enough of a clip."""
if not 0.0 <= minimum_span <= 1.0:
raise ValueError("minimum_span must be in [0, 1]")
if not modalities or any(name not in MODALITIES for name in modalities):
raise ValueError("modalities must be a non-empty subset of MODALITIES")
losses = []
for name in modalities:
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
span = mu[:, -1] - mu[:, 0]
losses.append(F.relu(minimum_span - span).square().mean())
return torch.stack(losses).mean()
def attention_diversity_loss(
output: AlignmentOutput,
modalities: tuple[str, ...] = MODALITIES,
min_separation: int = 6,
) -> Tensor:
"""Penalize similar attention rows for slots far apart on the grid."""
if min_separation < 1:
raise ValueError("min_separation must be at least one")
if not modalities or any(name not in MODALITIES for name in modalities):
raise ValueError("modalities must be a non-empty subset of MODALITIES")
return torch.stack(
[attention_row_similarity(output.weights[name], min_separation).mean() for name in modalities]
).mean()
def weak_temporal_band_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
targets: Mapping[str, Tensor],
margin: float = 0.1,
) -> Tensor:
"""Allow soft alignment while keeping expected times near weak slot anchors."""
if margin < 0:
raise ValueError("margin must be non-negative")
if not targets:
return next(iter(output.weights.values())).sum() * 0.0
losses = []
for name, target in targets.items():
if name not in MODALITIES:
raise ValueError(f"unknown modality in temporal targets: {name}")
mu = torch.bmm(output.weights[name], times[name].unsqueeze(-1)).squeeze(-1)
mu = mu / durations[:, None].clamp_min(torch.finfo(mu.dtype).eps)
if target.shape != mu.shape:
raise ValueError(f"band target for {name} must have shape {tuple(mu.shape)}")
distance = (mu - target).abs()
losses.append(F.relu(distance - margin).square().mean())
return torch.stack(losses).mean()
def alignment_training_loss(
output: AlignmentOutput,
times: Mapping[str, Tensor],
durations: Tensor,
reconstruction: Tensor,
*,
lambda_rec: float = 1.0,
lambda_con: float = 1.0,
lambda_mono: float = 0.1,
lambda_span: float = 0.0,
lambda_div: float = 0.0,
lambda_band: float = 0.0,
epsilon: float = 0.02,
minimum_span: float = 0.7,
coverage_modalities: tuple[str, ...] = MODALITIES,
diversity_modalities: tuple[str, ...] = MODALITIES,
diversity_min_separation: int = 6,
band_targets: Mapping[str, Tensor] | None = None,
band_margin: float = 0.1,
) -> tuple[Tensor, dict[str, Tensor]]:
"""Shared M3/M4 training objective; emotion labels are deliberately unused."""
contrastive = cross_modal_contrastive_loss(output.aligned)
monotonicity = temporal_monotonicity_loss(output, times, durations, epsilon)
span = temporal_span_loss(
output, times, durations, minimum_span, modalities=coverage_modalities
)
diversity = attention_diversity_loss(
output, diversity_modalities, min_separation=diversity_min_separation
)
band = weak_temporal_band_loss(
output, times, durations, band_targets or {}, margin=band_margin
)
total = (
lambda_rec * reconstruction
+ lambda_con * contrastive
+ lambda_mono * monotonicity
+ lambda_span * span
+ lambda_div * diversity
+ lambda_band * band
)
return total, {
"reconstruction": reconstruction,
"contrastive": contrastive,
"monotonicity": monotonicity,
"span": span,
"diversity": diversity,
"band": band,
"total": total,
}