147 lines
6.6 KiB
Python
147 lines
6.6 KiB
Python
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
|