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

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()