354 lines
15 KiB
Python
354 lines
15 KiB
Python
"""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()
|