176 lines
6.7 KiB
Python
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,
|
|
}
|