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

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