274 lines
9.7 KiB
Python
274 lines
9.7 KiB
Python
"""Single-clip test of explicit source-time identities in learned alignment."""
|
|
|
|
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,
|
|
SIGMA,
|
|
_example_arrays,
|
|
_gaussian_alignment_kl,
|
|
_gaussian_targets,
|
|
_gradient_norms,
|
|
_metric_rows,
|
|
_plot_example,
|
|
_seed_everything,
|
|
)
|
|
from .experiment_data import fit_feature_stats, load_feature_samples, collate_feature_samples
|
|
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
|
|
from .types import MODALITIES
|
|
|
|
|
|
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_encoding: 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_encoding,
|
|
)
|
|
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_encoding,
|
|
)
|
|
|
|
|
|
def _run_trial(
|
|
*,
|
|
method: str,
|
|
variant: str,
|
|
source_time_encoding: bool,
|
|
sample: Any,
|
|
stats: Any,
|
|
device: torch.device,
|
|
steps: int,
|
|
seed: int,
|
|
output_dir: Path,
|
|
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
|
|
_seed_everything(seed)
|
|
dimensions = {name: sample.features[name].shape[1] for name in MODALITIES}
|
|
model = _make_model(method, dimensions, source_time_encoding).to(device)
|
|
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.0)
|
|
sequences, durations, _ = collate_feature_samples([sample], stats, device)
|
|
history: list[dict[str, Any]] = []
|
|
print(
|
|
f"[D4 {variant}] sample={sample.sample_id} steps={steps} "
|
|
f"source_time_encoding={source_time_encoding} device={device}",
|
|
flush=True,
|
|
)
|
|
model.train()
|
|
for step in range(1, steps + 1):
|
|
output = model(sequences, durations) if source_time_encoding 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 D4 loss at {variant} step {step}")
|
|
row: dict[str, Any] = {
|
|
"experiment": "D4",
|
|
"method": method,
|
|
"variant": variant,
|
|
"step": step,
|
|
"L_align": float(loss.detach().item()),
|
|
"source_time_encoding": source_time_encoding,
|
|
}
|
|
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"[D4 {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}/"
|
|
f"{grad_z_text}",
|
|
flush=True,
|
|
)
|
|
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
_write_csv(output_dir / "history.csv", history)
|
|
model.eval()
|
|
with torch.no_grad():
|
|
output = model(sequences, durations) if source_time_encoding else model(sequences)
|
|
metric_rows = _metric_rows(
|
|
method,
|
|
"D4",
|
|
sample,
|
|
output,
|
|
sequences,
|
|
durations,
|
|
stats=stats,
|
|
device=device,
|
|
)
|
|
arrays = _example_arrays(sample, output, sequences, method)
|
|
for row in metric_rows:
|
|
row["variant"] = variant
|
|
row["source_time_encoding"] = source_time_encoding
|
|
_write_csv(output_dir / "metrics.csv", metric_rows)
|
|
np.savez_compressed(output_dir / "example_alignment.npz", **arrays)
|
|
_plot_example(output_dir / variant, sample, arrays, method, "D4")
|
|
checkpoint = output_dir / "checkpoint.pt"
|
|
torch.save(
|
|
{
|
|
"experiment": "D4",
|
|
"method": method,
|
|
"variant": variant,
|
|
"source_time_encoding": source_time_encoding,
|
|
"absolute_position_encoding": method == "M4",
|
|
"seed": seed,
|
|
"steps": steps,
|
|
"model_state_dict": model.state_dict(),
|
|
},
|
|
checkpoint,
|
|
)
|
|
summary = {
|
|
"experiment": "D4",
|
|
"method": method,
|
|
"variant": variant,
|
|
"source_time_encoding": source_time_encoding,
|
|
"final_training_kl": history[-1]["L_align"],
|
|
"checkpoint": str(checkpoint),
|
|
}
|
|
del model
|
|
if device.type == "cuda":
|
|
torch.cuda.empty_cache()
|
|
return summary, metric_rows
|
|
|
|
|
|
def run(args: argparse.Namespace) -> dict[str, Any]:
|
|
start = 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)
|
|
sample_map = {sample.sample_id: sample for sample in samples}
|
|
if EXAMPLE_ID not in sample_map:
|
|
raise ValueError(f"diagnostic sample is missing: {EXAMPLE_ID}")
|
|
sample = sample_map[EXAMPLE_ID]
|
|
stats = fit_feature_stats([sample])
|
|
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),
|
|
)
|
|
summaries: list[dict[str, Any]] = []
|
|
all_metrics: list[dict[str, Any]] = []
|
|
for method, variant, use_source_time in specifications:
|
|
summary, metrics = _run_trial(
|
|
method=method,
|
|
variant=variant,
|
|
source_time_encoding=use_source_time,
|
|
sample=sample,
|
|
stats=stats,
|
|
device=device,
|
|
steps=args.steps,
|
|
seed=args.seed,
|
|
output_dir=output_root / variant,
|
|
)
|
|
summaries.append(summary)
|
|
all_metrics.extend(metrics)
|
|
_write_csv(output_root / "summary.csv", summaries)
|
|
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
|
|
manifest = {
|
|
"created_utc": datetime.now(timezone.utc).isoformat(),
|
|
"experiment": "D4",
|
|
"sample": sample.sample_id,
|
|
"sample_count": 1,
|
|
"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,
|
|
"sigma_normalized_time": SIGMA,
|
|
"optimizer": "AdamW",
|
|
"learning_rate": LEARNING_RATE,
|
|
"steps_per_model": args.steps,
|
|
"source_position_encoding": "fixed Fourier time code added to source key only; value remains the projected real feature",
|
|
"query_position_encoding": "M3 adds Fourier code at forced text-time centers; M4 keeps its fixed absolute sinusoidal slot code and adds the matching Fourier code at uniform slot centers",
|
|
"modalities": ["audio", "vision"],
|
|
"variants": summaries,
|
|
"elapsed_seconds": time.time() - start,
|
|
"interpretation_limits": [
|
|
"This is a one-sample overfit diagnostic, not a held-out accuracy result.",
|
|
"The timestamp-derived Gaussian target is a weak temporal prior, not human alignment ground truth.",
|
|
"The D4 variants test source/query time identity only; content features and downstream outputs still require separate evaluation.",
|
|
],
|
|
}
|
|
(output_root / "run_manifest.json").write_text(
|
|
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
print(f"[D4 done] 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=1000)
|
|
parser.add_argument("--seed", type=int, default=42)
|
|
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(
|
|
"--output-dir", type=Path, default=project / "outputs/alignment_debug/source_time"
|
|
)
|
|
return parser
|
|
|
|
|
|
def main() -> None:
|
|
args = build_parser().parse_args()
|
|
run(args)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|