from __future__ import annotations import numpy as np import torch import torch.nn.functional as F from torch import Tensor from .types import MODALITIES def alignment_trajectory(weights: Tensor, times: Tensor, durations: Tensor) -> Tensor: """Return expected normalized source time at each common-grid position.""" if weights.ndim != 3 or times.shape != (weights.shape[0], weights.shape[2]): raise ValueError("weights [B,K,L] and times [B,L] must agree") if durations.shape != (weights.shape[0],): raise ValueError("durations must have shape [B]") expected_seconds = torch.bmm(weights, times.unsqueeze(-1)).squeeze(-1) return expected_seconds / durations[:, None].clamp_min(torch.finfo(expected_seconds.dtype).eps) def monotonicity_violation_rate(trajectory: Tensor, epsilon: float = 0.02) -> Tensor: """Per-sample fraction of adjacent grid pairs that move backwards by epsilon.""" if trajectory.ndim != 2 or trajectory.shape[1] < 2: raise ValueError("trajectory must have shape [B, K] with K >= 2") if epsilon < 0: raise ValueError("epsilon must be non-negative") return ((trajectory[:, :-1] - trajectory[:, 1:]) > epsilon).float().mean(dim=1) def normalized_attention_entropy(weights: Tensor, valid: Tensor) -> Tensor: """Per-row entropy normalized by the number of valid source positions.""" if weights.ndim != 3 or valid.shape != (weights.shape[0], weights.shape[2]): raise ValueError("weights [B,K,L] and valid [B,L] must agree") safe = weights.clamp_min(torch.finfo(weights.dtype).tiny) entropy = -(weights * safe.log()).sum(dim=-1) counts = valid.sum(dim=-1).clamp_min(1) denominator = counts.float().log().clamp_min(torch.finfo(torch.float32).eps) normalized = entropy / denominator[:, None] return torch.where(counts[:, None] > 1, normalized, torch.zeros_like(normalized)) def attention_row_similarity(weights: Tensor, min_separation: int = 1) -> Tensor: """Mean cosine similarity between attention rows, per sample. ``min_separation=1`` compares every distinct pair (the C_row collapse score). A value of 6 compares only pairs more than five slots apart, matching the training diversity loss. Scores near one mean that slots attend to nearly the same source distribution. """ if weights.ndim != 3: raise ValueError("weights must have shape [B, K, L]") if min_separation < 1: raise ValueError("min_separation must be at least one") batch_size, grid_size, _ = weights.shape if grid_size <= min_separation: return torch.zeros(batch_size, dtype=weights.dtype, device=weights.device) normalized = F.normalize(weights, p=2, dim=-1, eps=1e-12) similarities = torch.bmm(normalized, normalized.transpose(1, 2)) pair_mask = torch.triu( torch.ones(grid_size, grid_size, dtype=torch.bool, device=weights.device), diagonal=min_separation, ) return similarities[:, pair_mask].mean(dim=-1) def attention_width80(weights: Tensor, threshold: float = 0.8) -> Tensor: """Shortest contiguous source-index span containing the requested mass. Returns integer widths with shape ``[B, K]``. This is a diagnostic, not a score to maximize or minimize on its own. """ if weights.ndim != 3 or not 0 < threshold <= 1: raise ValueError("weights must be [B,K,L] and threshold in (0, 1]") rows = weights.detach().to(device="cpu", dtype=torch.float64).numpy() widths = np.empty(rows.shape[:2], dtype=np.int64) for batch_index in range(rows.shape[0]): for grid_index in range(rows.shape[1]): row = rows[batch_index, grid_index] left = 0 mass = 0.0 best = len(row) for right, value in enumerate(row): mass += float(value) while left <= right and mass - float(row[left]) >= threshold: mass -= float(row[left]) left += 1 if mass + 1e-12 >= threshold: best = min(best, right - left + 1) widths[batch_index, grid_index] = best return torch.from_numpy(widths) def retrieval_metrics(query: Tensor, target: Tensor, chunk_size: int = 256) -> dict[str, float]: """Grid-slot retrieval; same flattened sample/slot index is the positive. Use only as a representation-consistency probe. It is not independent temporal ground truth; report human-labeled temporal scores separately. """ if query.ndim != 3 or target.ndim != 3 or query.shape != target.shape: raise ValueError("query and target must have matching [N, K, D] shapes") if query.shape[0] * query.shape[1] < 1: raise ValueError("retrieval requires at least one grid position") q = F.normalize(query.flatten(0, 1), dim=-1) t = F.normalize(target.flatten(0, 1), dim=-1) total = q.shape[0] ranks = torch.empty(total, dtype=torch.long, device=q.device) for start in range(0, total, chunk_size): stop = min(start + chunk_size, total) scores = q[start:stop] @ t.T positives = scores[torch.arange(stop - start, device=q.device), torch.arange(start, stop, device=q.device)] ranks[start:stop] = 1 + (scores > positives[:, None]).sum(dim=1) ranks_f = ranks.float() return { "r_at_1": float((ranks <= 1).float().mean().item()), "r_at_5": float((ranks <= min(5, total)).float().mean().item()), "mrr": float((1.0 / ranks_f).mean().item()), "queries": float(total), } def summarize_alignment( weights: dict[str, Tensor], times: dict[str, Tensor], valid: dict[str, Tensor], durations: Tensor, epsilon: float = 0.02, ) -> dict[str, dict[str, float]]: """Produce sample-aggregated E1-E3 diagnostics for each modality.""" summary: dict[str, dict[str, float]] = {} for name in MODALITIES: trajectory = alignment_trajectory(weights[name], times[name], durations) mvr = monotonicity_violation_rate(trajectory, epsilon) entropy = normalized_attention_entropy(weights[name], valid[name]) width = attention_width80(weights[name]) summary[name] = { "mvr": float(mvr.mean().item()), "normalized_entropy": float(entropy.mean().item()), "width80_indices": float(width.float().mean().item()), "mean_time_start": float(trajectory[:, 0].mean().item()), "mean_time_end": float(trajectory[:, -1].mean().item()), "trajectory_span_fraction": float( (trajectory[:, -1] - trajectory[:, 0]).mean().item() ), } return summary