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