"""Evaluate TSFA attention maps using identical raw source content for every method. This probe pools the same train-fold-standardized BERT, audio, and DeiT features with each method's alignment matrix. It excludes native model value/output projections from the representation being scored. """ from __future__ import annotations import csv import json import shutil import time from datetime import datetime, timezone from pathlib import Path from typing import Any, Mapping, Sequence import numpy as np import torch from .correspondence_eval import _cluster_bootstrap, _write_csv from .experiment_data import ( FeatureSample, collate_feature_samples, fit_feature_stats, load_feature_samples, ) from .m4_shared_latent_eval import _bootstrap_summary from .tsfa_experiment import ( ALL_METHODS, BASELINE_VARIANTS, GRID_SIZE, TSFA_VARIANTS, _ablation_summary, _collect_fold_features, _content_summary, _evaluate_fixed_projector, _fit_method_probe, _generate_tsfa_outputs, _load_semantic_checkpoint, build_parser, ) from .types import MODALITIES def _pool_same_raw_features( samples: Sequence[FeatureSample], feature_stats: Any, weights_by_id: Mapping[str, Mapping[str, np.ndarray]], device: torch.device, batch_size: int, ) -> dict[str, dict[str, np.ndarray]]: content: dict[str, dict[str, np.ndarray]] = {} with torch.no_grad(): for start in range(0, len(samples), batch_size): batch_samples = list(samples[start : start + batch_size]) sequences, _, _ = collate_feature_samples(batch_samples, feature_stats, device) for index, sample in enumerate(batch_samples): content[sample.sample_id] = {} for modality in MODALITIES: length = len(sample.features[modality]) weights = torch.as_tensor( weights_by_id[sample.sample_id][modality], dtype=torch.float32, device=device, ) if weights.shape != (GRID_SIZE, length): raise ValueError(f"unexpected alignment shape for {sample.sample_id}/{modality}") source = sequences[modality].features[index, :length] content[sample.sample_id][modality] = ( (weights @ source).cpu().numpy().astype(np.float32, copy=False) ) return content def _paired_summary(rows: Sequence[Mapping[str, Any]], seed: int) -> list[dict[str, Any]]: by_key = {(row["method"], row["sample_id"]): row for row in rows} sample_ids = sorted({row["sample_id"] for row in rows}) contrasts = ( ("TSFA-main", "M4_sourceTime"), ("TSFA-main", "TSFA-random"), ("TSFA-main", "TSFA-global"), ("TSFA-multiply", "TSFA-main"), ("TSFA-main", "M3_noSourceTime"), ) output = [] for contrast_index, (left, right) in enumerate(contrasts): for metric_index, metric in enumerate(( "content_auc_mean", "canonical_pairwise_time_mae_mean" )): differences = [] for sample_id in sample_ids: left_row = by_key[(left, sample_id)] right_row = by_key[(right, sample_id)] if left_row["video_id"] != right_row["video_id"]: raise ValueError(f"video group mismatch for {sample_id}") differences.append({ "video_id": left_row["video_id"], "difference": float(left_row[metric]) - float(right_row[metric]), }) mean, low, high, groups = _cluster_bootstrap( differences, "difference", seed=seed + contrast_index * 101 + metric_index, repetitions=2000, ) output.append({ "left_method": left, "right_method": right, "metric": metric, "clip_count": len(differences), "video_id_count": groups, "left_minus_right_video_macro_mean": mean, "ci95_low": low, "ci95_high": high, }) return output def _temporal_diagnostic_rows( method: str, fold: int, sample: FeatureSample, weights: Mapping[str, np.ndarray], draw: int = 0, ) -> list[dict[str, Any]]: output = [] for modality in MODALITIES: valid = np.asarray(sample.valid[modality], dtype=bool) attention = np.asarray(weights[modality], dtype=np.float64)[:, valid] times = np.asarray(sample.times[modality], dtype=np.float64)[valid] / max(sample.duration_s, 1e-8) centers = attention @ times backward = centers[:-1] - centers[1:] entropy = -(attention * np.log(np.maximum(attention, 1e-12))).sum(axis=1) output.append({ "method": method, "fold": fold, "sample_id": sample.sample_id, "video_id": sample.group_id, "draw": draw, "modality": modality, "mvr_epsilon_0_01": float(np.mean(backward > 0.01)), "mvr_epsilon_0_02": float(np.mean(backward > 0.02)), "mvr_epsilon_0_05": float(np.mean(backward > 0.05)), "time_span_ratio": float(centers[-1] - centers[0]), "normalized_entropy_mean": float(entropy.mean() / max(np.log(attention.shape[1]), 1e-12)), "source_coverage_rate": float(np.mean(attention.sum(axis=0) > 1e-12)), }) return output def run(args: Any) -> None: started = time.time() device = torch.device( ("cuda" if torch.cuda.is_available() else "cpu") if args.device == "auto" else args.device ) if device.type == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA is unavailable") samples = load_feature_samples(args.feature_dir, args.manifest) samples_by_id = {sample.sample_id: sample for sample in samples} splits = json.loads(args.splits.read_text(encoding="utf-8")) if len(splits) != 5: raise ValueError("expected five grouped folds") store = torch.load(args.output_dir / "probe_checkpoints.pt", map_location="cpu", weights_only=False) content_rows: list[dict[str, Any]] = [] curve_rows: list[dict[str, Any]] = [] temporal_diagnostic_rows: list[dict[str, Any]] = [] history_rows: list[dict[str, Any]] = [] probe_store: dict[str, Any] = {} heldout_ids = [] for split in splits: fold = int(split["fold"]) train = [samples_by_id[sample_id] for sample_id in split["train_sample_ids"]] validation = [samples_by_id[sample_id] for sample_id in split["validation_sample_ids"]] if {sample.group_id for sample in train} & {sample.group_id for sample in validation}: raise ValueError(f"video_id leakage in fold {fold}") heldout_ids.extend(sample.sample_id for sample in validation) feature_stats = fit_feature_stats(train) _, baseline_weights, temporal_by_id = _collect_fold_features( fold=fold, train_samples=train, validation_samples=validation, feature_stats=feature_stats, checkpoint_root=args.checkpoint_root, device=device, batch_size=args.batch_size, ) branch = _load_semantic_checkpoint(store, fold, device) all_samples = [*train, *validation] all_ids = [sample.sample_id for sample in all_samples] fold_weights = {method: baseline_weights[method] for method in BASELINE_VARIANTS} for method in TSFA_VARIANTS: _, weights, _ = _generate_tsfa_outputs( method=method, fold=fold, sample_ids=all_ids, samples_by_id=samples_by_id, temporal_by_id=temporal_by_id, branch=branch, device=device, delta=args.delta, seed=args.seed, draw=0, batch_size=args.batch_size, ) fold_weights[method] = weights for method in ALL_METHODS: pooled = _pool_same_raw_features( all_samples, feature_stats, fold_weights[method], device, args.batch_size ) metrics, curves, _ = _fit_method_probe( method=method, fold=fold, train_samples=train, validation_samples=validation, content_by_id=pooled, device=device, args=args, history_rows=history_rows, checkpoint_store=probe_store, ) content_rows.extend(metrics) curve_rows.extend(curves) for sample in validation: temporal_diagnostic_rows.extend(_temporal_diagnostic_rows( method, fold, sample, fold_weights[method][sample.sample_id] )) if method == "TSFA-random": for draw in range(1, args.random_window_repeats): _, random_weights, _ = _generate_tsfa_outputs( method=method, fold=fold, sample_ids=[sample.sample_id for sample in validation], samples_by_id=samples_by_id, temporal_by_id=temporal_by_id, branch=branch, device=device, delta=args.delta, seed=args.seed, draw=draw, batch_size=args.batch_size, ) random_pooled = _pool_same_raw_features( validation, feature_stats, random_weights, device, args.batch_size ) repeated_metrics, repeated_curves, _ = _evaluate_fixed_projector( method=method, fold=fold, validation_samples=validation, content_by_id=random_pooled, projector_state=probe_store[ f"fold_{fold:02d}/{method}/correspondence_probe" ]["state_dict"], device=device, ) content_rows.extend({**row, "draw": draw} for row in repeated_metrics) curve_rows.extend({**row, "draw": draw} for row in repeated_curves) for sample in validation: temporal_diagnostic_rows.extend(_temporal_diagnostic_rows( method, fold, sample, random_weights[sample.sample_id], draw )) print(f"[TSFA alignment-only fold {fold}] heldout={len(validation)}", flush=True) del branch, baseline_weights, temporal_by_id, fold_weights if device.type == "cuda": torch.cuda.empty_cache() if len(heldout_ids) != 100 or len(set(heldout_ids)) != 100: raise ValueError("held-out fold coverage is not exactly 100 distinct samples") output_dir = args.output_dir content_summary, curve_summary = _content_summary(content_rows, curve_rows, args.seed + 901) with (output_dir / "temporal_metrics_by_clip.csv").open( newline="", encoding="utf-8-sig" ) as handle: temporal_rows = list(csv.DictReader(handle)) ablation_summary, ablation_by_clip = _ablation_summary( content_rows, temporal_rows, args.seed + 902 ) paired = _paired_summary(ablation_by_clip, args.seed + 903) temporal_diagnostic_summary = _bootstrap_summary( temporal_diagnostic_rows, ("method", "modality"), ( "mvr_epsilon_0_01", "mvr_epsilon_0_02", "mvr_epsilon_0_05", "time_span_ratio", "normalized_entropy_mean", "source_coverage_rate", ), seed=args.seed + 904, ) _write_csv(output_dir / "alignment_only_content_by_clip.csv", content_rows) _write_csv(output_dir / "alignment_only_content_summary.csv", content_summary) _write_csv(output_dir / "alignment_only_shift_curve_summary.csv", curve_summary) _write_csv(output_dir / "alignment_only_ablation_by_clip.csv", ablation_by_clip) _write_csv(output_dir / "alignment_only_ablation_summary.csv", ablation_summary) _write_csv(output_dir / "alignment_only_paired_contrasts.csv", paired) _write_csv(output_dir / "alignment_only_probe_training_history.csv", history_rows) _write_csv(output_dir / "tsfa_temporal_diagnostics_by_clip.csv", temporal_diagnostic_rows) _write_csv(output_dir / "tsfa_temporal_diagnostics_summary.csv", temporal_diagnostic_summary) torch.save(probe_store, output_dir / "alignment_only_probe_checkpoints.pt") run_manifest = { "created_utc": datetime.now(timezone.utc).isoformat(), "experiment": "TSFA shared raw-content alignment-only probe", "sample_count": len(samples), "heldout_count": len(heldout_ids), "video_id_count": len({sample.group_id for sample in samples}), "fold_count": len(splits), "seed": args.seed, "delta": args.delta, "probe_epochs": args.probe_epochs, "random_window_repeats": args.random_window_repeats, "feature_protocol": "Within each fold, feature normalization is fitted on 80 training clips. For every method and modality, the same normalized raw source feature matrix is pooled using that method's A^m; native M3/M4/TSFA value and output projections are excluded from scored features.", "probe_protocol": "One 64-dimensional linear projector per modality, trained on 80 training clips with the same within-clip InfoNCE protocol and identical fold seed across methods. Scores use only 20 held-out clips per fold.", "interpretation_limit": "A same-slot positive is a timestamp/slot convention, not independently annotated semantic ground truth. Source-time attention can still encode time through selected raw values.", "device": str(device), "elapsed_seconds": time.time() - started, } (output_dir / "alignment_only_run_manifest.json").write_text( json.dumps(run_manifest, ensure_ascii=False, indent=2), encoding="utf-8" ) bundle = output_dir / "report_bundle" for name in ( "alignment_only_content_summary.csv", "alignment_only_ablation_summary.csv", "alignment_only_paired_contrasts.csv", "alignment_only_run_manifest.json", "tsfa_temporal_diagnostics_summary.csv", ): shutil.copy2(output_dir / name, bundle / name) bundle_readme = bundle / "README.md" note = ( "\nThe `alignment_only_*` summaries pool identical train-fold-standardized " "raw source features through each method's alignment matrix. They isolate " "source selection from native model value/output projections.\n" ) current = bundle_readme.read_text(encoding="utf-8") if "The `alignment_only_*` summaries" not in current: bundle_readme.write_text(current + note, encoding="utf-8") print( f"[TSFA alignment-only complete] samples={len(heldout_ids)} " f"elapsed={run_manifest['elapsed_seconds']:.1f}s output={output_dir}", flush=True, ) def main() -> None: args = build_parser().parse_args() if args.finalize_existing: raise ValueError("--finalize-existing belongs to q1.tsfa_experiment") if not 0 < args.delta <= 1: raise ValueError("--delta must be in (0,1]") run(args) if __name__ == "__main__": main()