建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,353 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user