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

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