383 lines
15 KiB
Python
383 lines
15 KiB
Python
"""Grouped held-out check for learned source-time positional features."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import csv
|
|
import json
|
|
import platform
|
|
import time
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
import torch
|
|
from torch import nn
|
|
|
|
from .alignment_debug import (
|
|
EXAMPLE_ID,
|
|
GRID_SIZE,
|
|
HEADS,
|
|
HIDDEN_SIZE,
|
|
LEARNING_RATE,
|
|
_example_arrays,
|
|
_gaussian_alignment_kl,
|
|
_gaussian_targets,
|
|
_gradient_norms,
|
|
_metric_rows,
|
|
_plot_example,
|
|
_seed_everything,
|
|
)
|
|
from .experiment_data import (
|
|
FeatureSample,
|
|
FeatureStats,
|
|
collate_feature_samples,
|
|
fit_feature_stats,
|
|
load_feature_samples,
|
|
)
|
|
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
|
|
from .types import MODALITIES
|
|
|
|
|
|
BATCH_SIZE = 8
|
|
DEFAULT_STEPS = 500
|
|
|
|
|
|
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
|
if not rows:
|
|
return
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
fields = list(dict.fromkeys(key for row in rows for key in row))
|
|
with path.open("w", newline="", encoding="utf-8-sig") as handle:
|
|
writer = csv.DictWriter(handle, fieldnames=fields)
|
|
writer.writeheader()
|
|
writer.writerows(rows)
|
|
|
|
|
|
def _make_model(method: str, dimensions: dict[str, int], source_time: bool) -> nn.Module:
|
|
if method == "M3":
|
|
return TextAnchoredCrossAttention(
|
|
dimensions,
|
|
grid_size=GRID_SIZE,
|
|
hidden_size=HIDDEN_SIZE,
|
|
heads=HEADS,
|
|
dropout=0.0,
|
|
source_time_encoding=source_time,
|
|
)
|
|
return SharedLatentTimeline(
|
|
dimensions,
|
|
grid_size=GRID_SIZE,
|
|
hidden_size=HIDDEN_SIZE,
|
|
heads=HEADS,
|
|
dropout=0.0,
|
|
absolute_position_encoding=True,
|
|
source_time_encoding=source_time,
|
|
)
|
|
|
|
|
|
def _batches(samples: list[FeatureSample], rng: np.random.Generator):
|
|
order = rng.permutation(len(samples)).tolist()
|
|
for start in range(0, len(order), BATCH_SIZE):
|
|
yield [samples[index] for index in order[start : start + BATCH_SIZE]]
|
|
|
|
|
|
def _train_one(
|
|
*,
|
|
method: str,
|
|
variant: str,
|
|
source_time: bool,
|
|
train_samples: list[FeatureSample],
|
|
validation_samples: list[FeatureSample],
|
|
stats: FeatureStats,
|
|
example_id: str,
|
|
output_dir: Path,
|
|
device: torch.device,
|
|
steps: int,
|
|
seed: int,
|
|
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
|
|
_seed_everything(seed)
|
|
dimensions = {
|
|
name: train_samples[0].features[name].shape[1] for name in MODALITIES
|
|
}
|
|
model = _make_model(method, dimensions, source_time).to(device)
|
|
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0)
|
|
rng = np.random.default_rng(seed)
|
|
history: list[dict[str, Any]] = []
|
|
print(
|
|
f"[D5 {variant}] train={len(train_samples)} heldout={len(validation_samples)} "
|
|
f"groups={len({s.group_id for s in train_samples})}/"
|
|
f"{len({s.group_id for s in validation_samples})} steps={steps}",
|
|
flush=True,
|
|
)
|
|
model.train()
|
|
for step in range(1, steps + 1):
|
|
batch_indices = rng.choice(
|
|
len(train_samples), size=min(BATCH_SIZE, len(train_samples)), replace=False
|
|
)
|
|
batch_samples = [train_samples[int(index)] for index in batch_indices]
|
|
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
|
|
output = model(sequences, durations) if source_time else model(sequences)
|
|
targets = _gaussian_targets(method, output, sequences, durations)
|
|
loss = _gaussian_alignment_kl(output, targets)
|
|
if not torch.isfinite(loss):
|
|
raise FloatingPointError(f"non-finite D5 loss for {variant} at step {step}")
|
|
row: dict[str, Any] = {
|
|
"experiment": "D5",
|
|
"method": method,
|
|
"variant": variant,
|
|
"step": step,
|
|
"L_align": float(loss.detach().item()),
|
|
"source_time_encoding": source_time,
|
|
}
|
|
if step == 1 or step % 20 == 0 or step == steps:
|
|
row.update(_gradient_norms(model, method, {"align": loss}, ("align",)))
|
|
optimizer.zero_grad(set_to_none=True)
|
|
loss.backward()
|
|
nn.utils.clip_grad_norm_(model.parameters(), 2.0)
|
|
optimizer.step()
|
|
history.append(row)
|
|
if step == 1 or step % 100 == 0 or step == steps:
|
|
grad_z = row.get("grad_align_Z")
|
|
grad_z_text = f"{grad_z:.3g}" if grad_z is not None else "NA"
|
|
print(
|
|
f"[D5 {variant} {step}/{steps}] KL={row['L_align']:.5f} "
|
|
f"grad_Q/K/Z={row.get('grad_align_WQ', 0):.3g}/"
|
|
f"{row.get('grad_align_WK', 0):.3g}/{grad_z_text}",
|
|
flush=True,
|
|
)
|
|
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
_write_csv(output_dir / "history.csv", history)
|
|
model.eval()
|
|
metric_rows: list[dict[str, Any]] = []
|
|
heldout_arrays: dict[str, np.ndarray] | None = None
|
|
with torch.no_grad():
|
|
eval_rng = np.random.default_rng(0)
|
|
for batch_samples in _batches(validation_samples, eval_rng):
|
|
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
|
|
output = model(sequences, durations) if source_time else model(sequences)
|
|
for index, sample in enumerate(batch_samples):
|
|
one_sequences = {
|
|
name: type(sequences[name])(
|
|
features=sequences[name].features[index : index + 1],
|
|
times=sequences[name].times[index : index + 1],
|
|
valid=sequences[name].valid[index : index + 1],
|
|
)
|
|
for name in MODALITIES
|
|
}
|
|
one_output = type(output)(
|
|
weights={name: output.weights[name][index : index + 1] for name in MODALITIES},
|
|
aligned={name: output.aligned[name][index : index + 1] for name in MODALITIES},
|
|
fallback_rows=output.fallback_rows,
|
|
)
|
|
rows = _metric_rows(
|
|
method,
|
|
"D5",
|
|
sample,
|
|
one_output,
|
|
one_sequences,
|
|
durations[index : index + 1],
|
|
stats=stats,
|
|
device=device,
|
|
)
|
|
for metric_row in rows:
|
|
metric_row["variant"] = variant
|
|
metric_row["source_time_encoding"] = source_time
|
|
metric_rows.extend(rows)
|
|
if sample.sample_id == example_id:
|
|
heldout_arrays = _example_arrays(
|
|
sample, one_output, one_sequences, method
|
|
)
|
|
|
|
_write_csv(output_dir / "heldout_metrics.csv", metric_rows)
|
|
if heldout_arrays is not None:
|
|
heldout_sample = next(s for s in validation_samples if s.sample_id == example_id)
|
|
np.savez_compressed(output_dir / "heldout_alignment.npz", **heldout_arrays)
|
|
_plot_example(output_dir / variant, heldout_sample, heldout_arrays, method, "D5 held-out")
|
|
checkpoint = output_dir / "checkpoint.pt"
|
|
torch.save(
|
|
{
|
|
"experiment": "D5",
|
|
"method": method,
|
|
"variant": variant,
|
|
"source_time_encoding": source_time,
|
|
"absolute_position_encoding": method == "M4",
|
|
"seed": seed,
|
|
"steps": steps,
|
|
"train_sample_ids": [sample.sample_id for sample in train_samples],
|
|
"validation_sample_ids": [sample.sample_id for sample in validation_samples],
|
|
"model_state_dict": model.state_dict(),
|
|
},
|
|
checkpoint,
|
|
)
|
|
training_summary = {
|
|
"experiment": "D5",
|
|
"method": method,
|
|
"variant": variant,
|
|
"source_time_encoding": source_time,
|
|
"final_training_kl": history[-1]["L_align"],
|
|
"heldout_sample_count": len(validation_samples),
|
|
"checkpoint": str(checkpoint),
|
|
}
|
|
del model
|
|
if device.type == "cuda":
|
|
torch.cuda.empty_cache()
|
|
return metric_rows, history, training_summary
|
|
|
|
|
|
def run(args: argparse.Namespace) -> dict[str, Any]:
|
|
started = time.time()
|
|
if args.device == "auto":
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
else:
|
|
device = torch.device(args.device)
|
|
if device.type == "cuda" and not torch.cuda.is_available():
|
|
raise RuntimeError("CUDA was requested but is unavailable")
|
|
samples = load_feature_samples(args.feature_dir, args.manifest)
|
|
by_id = {sample.sample_id: sample for sample in samples}
|
|
with args.splits.open("r", encoding="utf-8-sig") as handle:
|
|
folds = json.load(handle)
|
|
fold = next((item for item in folds if item["fold"] == args.fold), None)
|
|
if fold is None:
|
|
raise ValueError(f"fold {args.fold} is not present in {args.splits}")
|
|
train_samples = [by_id[sample_id] for sample_id in fold["train_sample_ids"]]
|
|
validation_samples = [by_id[sample_id] for sample_id in fold["validation_sample_ids"]]
|
|
train_groups = {sample.group_id for sample in train_samples}
|
|
validation_groups = {sample.group_id for sample in validation_samples}
|
|
if train_groups & validation_groups:
|
|
raise ValueError("train/validation video_id groups overlap")
|
|
example_id = args.example_id or validation_samples[0].sample_id
|
|
if example_id not in {sample.sample_id for sample in validation_samples}:
|
|
raise ValueError(f"held-out example is not in validation fold: {example_id}")
|
|
stats = fit_feature_stats(train_samples)
|
|
output_root = args.output_dir
|
|
output_root.mkdir(parents=True, exist_ok=True)
|
|
|
|
specifications = (
|
|
("M3", "M3_noSourceTime", False),
|
|
("M3", "M3_sourceTime", True),
|
|
("M4", "M4_noSourceTime", False),
|
|
("M4", "M4_sourceTime", True),
|
|
)
|
|
all_metrics: list[dict[str, Any]] = []
|
|
all_history: list[dict[str, Any]] = []
|
|
summaries: list[dict[str, Any]] = []
|
|
for method, variant, source_time in specifications:
|
|
metrics, history, summary = _train_one(
|
|
method=method,
|
|
variant=variant,
|
|
source_time=source_time,
|
|
train_samples=train_samples,
|
|
validation_samples=validation_samples,
|
|
stats=stats,
|
|
example_id=example_id,
|
|
output_dir=output_root / variant,
|
|
device=device,
|
|
steps=args.steps,
|
|
seed=args.seed,
|
|
)
|
|
all_metrics.extend(metrics)
|
|
all_history.extend(history)
|
|
summaries.append(summary)
|
|
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
|
|
_write_csv(output_root / "training_history.csv", all_history)
|
|
|
|
aggregate_rows: list[dict[str, Any]] = []
|
|
metric_names = (
|
|
"mvr",
|
|
"normalized_entropy",
|
|
"c_row",
|
|
"trajectory_span",
|
|
"mean_absolute_time_center_error",
|
|
"gaussian_target_kl",
|
|
)
|
|
for variant, modality in sorted(
|
|
{(row["variant"], row["modality"]) for row in all_metrics}
|
|
):
|
|
rows = [row for row in all_metrics if row["variant"] == variant and row["modality"] == modality]
|
|
aggregate: dict[str, Any] = {
|
|
"variant": variant,
|
|
"modality": modality,
|
|
"sample_count": len(rows),
|
|
}
|
|
for name in metric_names:
|
|
values = [float(row[name]) for row in rows if row.get(name) not in (None, "")]
|
|
aggregate[f"mean_{name}"] = float(np.mean(values)) if values else ""
|
|
aggregate_rows.append(aggregate)
|
|
_write_csv(output_root / "heldout_summary.csv", aggregate_rows)
|
|
_write_csv(output_root / "training_summary.csv", summaries)
|
|
|
|
manifest = {
|
|
"created_utc": datetime.now(timezone.utc).isoformat(),
|
|
"experiment": "D5",
|
|
"fold": args.fold,
|
|
"heldout_example": example_id,
|
|
"train_sample_count": len(train_samples),
|
|
"heldout_sample_count": len(validation_samples),
|
|
"train_video_ids": sorted(train_groups),
|
|
"heldout_video_ids": sorted(validation_groups),
|
|
"video_id_overlap": sorted(train_groups & validation_groups),
|
|
"seed": args.seed,
|
|
"device": str(device),
|
|
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
|
|
"python": platform.python_version(),
|
|
"torch": torch.__version__,
|
|
"grid_size": GRID_SIZE,
|
|
"optimizer": "AdamW",
|
|
"learning_rate": LEARNING_RATE,
|
|
"steps_per_model": args.steps,
|
|
"batch_size": BATCH_SIZE,
|
|
"loss": "timestamp-derived Gaussian target KL only",
|
|
"source_position_encoding": "fixed Fourier time code added to source key only; value remains projected content",
|
|
"query_position_encoding": "M3 adds Fourier code at text-time centers; M4 uses fixed absolute sinusoidal slots plus the matching Fourier code at uniform slot centers",
|
|
"normalization_fit_on_train_only": True,
|
|
"variants": summaries,
|
|
"elapsed_seconds": time.time() - started,
|
|
"interpretation_limits": [
|
|
"This is one grouped video_id split and one seed; it is a focused held-out diagnostic, not a final method ranking.",
|
|
"The Gaussian timestamp target is a weak temporal prior, not human alignment ground truth.",
|
|
"The target supplies approximate time location; this experiment tests transfer of the time-conditioned attention mechanism, not semantic correctness by itself.",
|
|
],
|
|
}
|
|
(output_root / "run_manifest.json").write_text(
|
|
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
print(
|
|
f"[D5 done] train={len(train_samples)} heldout={len(validation_samples)} "
|
|
f"video_id_groups={len(train_groups)}/{len(validation_groups)} "
|
|
f"elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}",
|
|
flush=True,
|
|
)
|
|
return manifest
|
|
|
|
|
|
def build_parser() -> argparse.ArgumentParser:
|
|
project = Path(__file__).resolve().parents[1]
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--steps", type=int, default=DEFAULT_STEPS)
|
|
parser.add_argument("--seed", type=int, default=42)
|
|
parser.add_argument("--fold", type=int, default=1)
|
|
parser.add_argument("--example-id", type=str, default=None)
|
|
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
|
parser.add_argument(
|
|
"--feature-dir", type=Path, default=project / "outputs/q1_features/features"
|
|
)
|
|
parser.add_argument("--manifest", type=Path, default=project / "outputs/audit/manifest.csv")
|
|
parser.add_argument(
|
|
"--splits", type=Path, default=project / "outputs/method_comparison/splits.json"
|
|
)
|
|
parser.add_argument(
|
|
"--output-dir", type=Path, default=project / "outputs/alignment_debug/heldout"
|
|
)
|
|
return parser
|
|
|
|
|
|
def main() -> None:
|
|
args = build_parser().parse_args()
|
|
run(args)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|