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, }