建立分批同步基线(基础文件)

This commit is contained in:
2026-09-23 23:24:01 +08:00
commit 7fc76aaafd
70 changed files with 18635 additions and 0 deletions
+776
View File
@@ -0,0 +1,776 @@
from __future__ import annotations
import argparse
import csv
import json
import math
import platform
import random
import time
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping, Sequence
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from .alignment import make_block_mask
from .compare_methods import AlignmentReconstructor, _write_csv
from .experiment_data import (
FeatureSample,
FeatureStats,
collate_feature_samples,
fit_feature_stats,
load_feature_samples,
)
from .losses import (
cross_modal_contrastive_loss,
temporal_span_loss,
weak_temporal_band_loss,
)
from .metrics import (
alignment_trajectory,
attention_row_similarity,
normalized_attention_entropy,
monotonicity_violation_rate,
)
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
from .types import MODALITIES
EXAMPLE_ID = "-tPCytz4rww/12"
SIGMA = 0.10
GRID_SIZE = 50
HIDDEN_SIZE = 128
HEADS = 4
SINGLE_SAMPLE_STEPS = 1000
FULL_DATA_STEPS = 500
BATCH_SIZE = 8
LEARNING_RATE = 1e-3
LOG_INTERVAL = 20
def _seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def _model(
method: str,
dimensions: Mapping[str, int],
*,
absolute_position_encoding: bool = False,
) -> nn.Module:
if method == "M3":
return TextAnchoredCrossAttention(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
)
if method == "M4":
return SharedLatentTimeline(
dimensions,
grid_size=GRID_SIZE,
hidden_size=HIDDEN_SIZE,
heads=HEADS,
dropout=0.0,
absolute_position_encoding=absolute_position_encoding,
)
raise ValueError(f"unknown method: {method}")
def _centers(
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
) -> tuple[dict[str, Tensor], Tensor]:
batch_size = durations.shape[0]
if method == "M3":
text_centers = torch.bmm(
output.weights["text"], sequences["text"].times.unsqueeze(-1)
).squeeze(-1)
text_centers = text_centers / durations[:, None].clamp_min(1e-8)
return {"audio": text_centers, "vision": text_centers}, text_centers
centers = (
torch.arange(GRID_SIZE, dtype=durations.dtype, device=durations.device) + 0.5
) / GRID_SIZE
centers = centers.unsqueeze(0).expand(batch_size, -1)
return {name: centers for name in MODALITIES}, centers
def _gaussian_targets(
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
) -> dict[str, Tensor]:
centers, _ = _centers(method, output, sequences, durations)
target_names = ("audio", "vision") if method == "M3" else MODALITIES
targets: dict[str, Tensor] = {}
for name in target_names:
times = sequences[name].times / durations[:, None].clamp_min(1e-8)
difference = (times[:, None, :] - centers[name][:, :, None]) / SIGMA
logits = -0.5 * difference.square()
logits = logits.masked_fill(~sequences[name].valid[:, None, :], -torch.inf)
targets[name] = torch.softmax(logits, dim=-1)
return targets
def _gaussian_alignment_kl(output: Any, targets: Mapping[str, Tensor]) -> Tensor:
losses = []
for name, target in targets.items():
predicted = output.weights[name].clamp_min(1e-8)
safe_target = target.clamp_min(1e-12)
row_kl = (safe_target * (safe_target.log() - predicted.log())).sum(dim=-1)
losses.append(row_kl.mean())
return torch.stack(losses).mean()
def _training_components(
experiment: str,
method: str,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
*,
decoder: AlignmentReconstructor | None,
block_generator: torch.Generator,
) -> tuple[Tensor, dict[str, Tensor], tuple[str, ...]]:
centers, _ = _centers(method, output, sequences, durations)
coverage_modalities = ("audio", "vision") if method == "M3" else MODALITIES
span = temporal_span_loss(
output,
{name: sequences[name].times for name in MODALITIES},
durations,
minimum_span=0.7,
modalities=coverage_modalities,
)
band = weak_temporal_band_loss(
output,
{name: sequences[name].times for name in MODALITIES},
durations,
centers,
margin=0.1,
)
targets = _gaussian_targets(method, output, sequences, durations)
align = _gaussian_alignment_kl(output, targets)
reconstruction = align.new_zeros(())
contrastive = align.new_zeros(())
if experiment == "D0":
total = 5.0 * span + 10.0 * band
active = ("span", "band", "total")
elif experiment in {"D1", "D2"}:
total = align
active = ("align", "total")
elif experiment == "D3":
if decoder is None:
raise ValueError("D3 requires the reconstruction decoder")
target_losses = []
batch_size, grid_size = output.aligned["text"].shape[:2]
for target in MODALITIES:
mask = make_block_mask(
batch_size,
grid_size,
0.2,
output.aligned[target].device,
generator=block_generator,
)
prediction = decoder(target, output.aligned, mask)
target_losses.append(
F.smooth_l1_loss(prediction[mask], output.aligned[target][mask])
)
reconstruction = torch.stack(target_losses).mean()
contrastive = cross_modal_contrastive_loss(output.aligned)
total = align + reconstruction + contrastive
active = ("align", "reconstruction", "contrastive", "total")
else:
raise ValueError(f"unknown experiment: {experiment}")
return total, {
"reconstruction": reconstruction,
"contrastive": contrastive,
"span": span,
"band": band,
"align": align,
"total": total,
}, active
def _attention_projection_parameters(model: nn.Module, method: str) -> tuple[list[nn.Parameter], nn.Parameter | None]:
if method == "M3":
layers = (model.audio_attention.attention, model.vision_attention.attention)
slot_parameter = None
else:
layers = tuple(model.attention[name].attention for name in MODALITIES)
slot_parameter = model.slots
parameters = [layer.in_proj_weight for layer in layers]
return parameters, slot_parameter
def _gradient_norms(
model: nn.Module,
method: str,
component_losses: Mapping[str, Tensor],
active_components: Sequence[str],
) -> dict[str, float | None]:
qk_parameters, slot_parameter = _attention_projection_parameters(model, method)
parameters = [*qk_parameters]
if slot_parameter is not None:
parameters.append(slot_parameter)
hidden_size = HIDDEN_SIZE
values: dict[str, float | None] = {}
for component in active_components:
loss = component_losses[component]
gradients = torch.autograd.grad(
loss,
parameters,
retain_graph=True,
allow_unused=True,
)
q_sq = torch.zeros((), device=loss.device)
k_sq = torch.zeros((), device=loss.device)
for gradient in gradients[: len(qk_parameters)]:
if gradient is None:
continue
q_sq = q_sq + gradient[:hidden_size].square().sum()
k_sq = k_sq + gradient[hidden_size : 2 * hidden_size].square().sum()
values[f"grad_{component}_WQ"] = float(q_sq.sqrt().item())
values[f"grad_{component}_WK"] = float(k_sq.sqrt().item())
if slot_parameter is not None:
slot_gradient = gradients[-1]
values[f"grad_{component}_Z"] = (
float(slot_gradient.norm().item()) if slot_gradient is not None else 0.0
)
else:
values[f"grad_{component}_Z"] = None
return values
def _metric_rows(
method: str,
experiment: str,
sample: FeatureSample,
output: Any,
sequences: Mapping[str, Any],
durations: Tensor,
*,
stats: FeatureStats,
device: torch.device,
) -> list[dict[str, Any]]:
targets = _gaussian_targets(method, output, sequences, durations)
rows = []
for name in MODALITIES:
weights = output.weights[name]
valid = sequences[name].valid
trajectory = alignment_trajectory(weights, sequences[name].times, durations)
entropy = normalized_attention_entropy(weights, valid)
centers, _ = _centers(method, output, sequences, durations)
row: dict[str, Any] = {
"experiment": experiment,
"method": method,
"sample_id": sample.sample_id,
"modality": name,
"mvr": float(monotonicity_violation_rate(trajectory).mean().item()),
"normalized_entropy": float(entropy.mean().item()),
"c_row": float(attention_row_similarity(weights).mean().item()),
"c_far": float(attention_row_similarity(weights, min_separation=6).mean().item()),
"expected_time_start": float(trajectory[0, 0].item()),
"expected_time_end": float(trajectory[0, -1].item()),
"trajectory_span": float((trajectory[0, -1] - trajectory[0, 0]).item()),
"mean_absolute_time_center_error": float(
(trajectory - centers.get(name, trajectory.new_full(trajectory.shape, float("nan"))))
.abs()
.mean()
.item()
)
if name in centers
else None,
"gaussian_target_kl": None,
}
if name in targets:
target = targets[name].clamp_min(1e-12)
prediction = weights.clamp_min(1e-8)
row["gaussian_target_kl"] = float(
(target * (target.log() - prediction.log())).sum(dim=-1).mean().item()
)
rows.append(row)
return rows
def _example_arrays(sample: FeatureSample, output: Any, sequences: Mapping[str, Any], method: str) -> dict[str, np.ndarray]:
targets = _gaussian_targets(method, output, sequences, torch.tensor([sample.duration_s], device=output.weights["text"].device))
arrays: dict[str, np.ndarray] = {"sample_id": np.asarray(sample.sample_id)}
for name in MODALITIES:
length = len(sample.times[name])
arrays[f"weights_{name}"] = output.weights[name][0, :, :length].detach().cpu().numpy()
arrays[f"times_{name}_s"] = sample.times[name].astype(np.float32, copy=False)
arrays[f"valid_{name}"] = sample.valid[name]
trajectory = (
output.weights[name][0, :, :length]
@ torch.as_tensor(sample.times[name], dtype=torch.float32, device=output.weights[name].device)
/ sample.duration_s
)
arrays[f"trajectory_{name}"] = trajectory.detach().cpu().numpy()
if name in targets:
arrays[f"target_{name}"] = targets[name][0, :, :length].detach().cpu().numpy()
return arrays
def _plot_example(
path_prefix: Path,
sample: FeatureSample,
arrays: Mapping[str, np.ndarray],
method: str,
experiment: str,
) -> None:
modalities = ("audio", "vision") if method == "M3" else MODALITIES
fig, axes = plt.subplots(len(modalities), 2, figsize=(12, 4 * len(modalities)), constrained_layout=True)
if len(modalities) == 1:
axes = np.asarray([axes])
for row, name in enumerate(modalities):
times = arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8)
extent = (float(times[0]), float(times[-1]), 0.0, 1.0)
prediction = arrays[f"weights_{name}"]
image = axes[row, 0].imshow(
prediction,
origin="lower",
aspect="auto",
interpolation="nearest",
extent=extent,
cmap="magma",
)
axes[row, 0].set_title(f"{name}: learned A")
axes[row, 0].set_xlabel("source time / clip duration")
axes[row, 0].set_ylabel("slot index / K")
fig.colorbar(image, ax=axes[row, 0], fraction=0.046, pad=0.04)
target_key = f"target_{name}"
if target_key in arrays:
target = arrays[target_key]
image_target = axes[row, 1].imshow(
target,
origin="lower",
aspect="auto",
interpolation="nearest",
extent=extent,
cmap="magma",
)
axes[row, 1].set_title(f"{name}: Gaussian target P")
fig.colorbar(image_target, ax=axes[row, 1], fraction=0.046, pad=0.04)
else:
axes[row, 1].imshow(
np.zeros_like(prediction),
origin="lower",
aspect="auto",
extent=extent,
cmap="magma",
)
axes[row, 1].set_title(f"{name}: target not used by M3")
axes[row, 1].set_xlabel("source time / clip duration")
axes[row, 1].set_ylabel("slot index / K")
fig.suptitle(f"{experiment} {method} · {sample.sample_id}")
fig.savefig(path_prefix.with_name(path_prefix.name + "_heatmap.png"), dpi=160)
plt.close(fig)
fig, ax = plt.subplots(figsize=(8, 5), constrained_layout=True)
x = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE
for name in MODALITIES:
trajectory = arrays[f"trajectory_{name}"]
ax.plot(x, trajectory, label=f"{name} actual")
if f"target_{name}" in arrays:
target_weights = arrays[f"target_{name}"]
target_time = target_weights @ arrays[f"times_{name}_s"] / max(sample.duration_s, 1e-8)
ax.plot(x, target_time, linestyle="--", alpha=0.7, label=f"{name} target")
ax.plot([0, 1], [0, 1], color="black", linestyle=":", alpha=0.6, label="uniform-time reference")
ax.set(xlim=(0, 1), ylim=(0, 1), xlabel="shared slot position", ylabel="expected normalized source time")
ax.grid(alpha=0.2)
ax.legend(fontsize=8, ncol=2)
ax.set_title(f"{experiment} {method} trajectory · {sample.sample_id}")
fig.savefig(path_prefix.with_name(path_prefix.name + "_trajectory.png"), dpi=160)
plt.close(fig)
def _batches(samples: Sequence[FeatureSample], batch_size: int, 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 _evaluate_samples(
method: str,
experiment: str,
model: nn.Module,
samples: Sequence[FeatureSample],
stats: FeatureStats,
*,
device: torch.device,
batch_size: int,
) -> tuple[list[dict[str, Any]], dict[str, np.ndarray] | None]:
rows: list[dict[str, Any]] = []
example_arrays: dict[str, np.ndarray] | None = None
model.eval()
rng = np.random.default_rng(0)
with torch.no_grad():
for batch_samples in _batches(samples, batch_size, rng):
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = 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,
)
one_duration = durations[index : index + 1]
rows.extend(
_metric_rows(
method,
experiment,
sample,
one_output,
one_sequences,
one_duration,
stats=stats,
device=device,
)
)
if sample.sample_id == EXAMPLE_ID:
example_arrays = _example_arrays(sample, one_output, one_sequences, method)
return rows, example_arrays
def _run_trial(
experiment: str,
method: str,
tag: str,
train_samples: Sequence[FeatureSample],
eval_samples: Sequence[FeatureSample],
stats: FeatureStats,
output_dir: Path,
*,
device: torch.device,
seed: int,
steps: int,
absolute_position_encoding: bool,
batch_size: int,
include_content_losses: bool,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
_seed_everything(seed)
dimensions = {name: train_samples[0].features[name].shape[1] for name in MODALITIES}
model = _model(method, dimensions, absolute_position_encoding=absolute_position_encoding).to(device)
decoder = AlignmentReconstructor(HIDDEN_SIZE, dropout=0.0).to(device) if include_content_losses else None
parameters = list(model.parameters()) + (list(decoder.parameters()) if decoder is not None else [])
optimizer = torch.optim.AdamW(parameters, lr=LEARNING_RATE, weight_decay=0.0)
block_generator = torch.Generator(device=device)
block_generator.manual_seed(seed + 73)
rng = np.random.default_rng(seed)
order = np.arange(len(train_samples))
cursor = 0
history: list[dict[str, Any]] = []
log_every = LOG_INTERVAL
active_names: tuple[str, ...] | None = None
print(
f"[{experiment} {tag}] samples={len(train_samples)} steps={steps} "
f"sinusoidal_PE={absolute_position_encoding} lr={LEARNING_RATE}",
flush=True,
)
model.train()
if decoder is not None:
decoder.train()
for step in range(1, steps + 1):
if len(train_samples) == 1:
batch_samples = [train_samples[0]]
else:
if cursor + batch_size > len(order):
order = rng.permutation(len(train_samples))
cursor = 0
indices = order[cursor : cursor + batch_size]
cursor += len(indices)
batch_samples = [train_samples[int(index)] for index in indices]
sequences, durations, _ = collate_feature_samples(batch_samples, stats, device)
output = model(sequences)
total, losses, active_components = _training_components(
experiment,
method,
output,
sequences,
durations,
decoder=decoder,
block_generator=block_generator,
)
active_names = active_components
if not torch.isfinite(total):
raise FloatingPointError(f"non-finite {experiment}/{tag} loss at step {step}")
row: dict[str, Any] = {
"experiment": experiment,
"method": method,
"variant": tag,
"step": step,
"epoch": math.ceil(step / max(1, math.ceil(len(train_samples) / batch_size))),
"learning_rate": LEARNING_RATE,
"absolute_position_encoding": absolute_position_encoding,
"L_rec": float(losses["reconstruction"].detach().item()),
"L_con": float(losses["contrastive"].detach().item()),
"L_span": float(losses["span"].detach().item()),
"L_band": float(losses["band"].detach().item()),
"L_align": float(losses["align"].detach().item()),
"L_total": float(total.detach().item()),
}
if step == 1 or step % log_every == 0 or step == steps:
row.update(_gradient_norms(model, method, losses, active_components))
optimizer.zero_grad(set_to_none=True)
total.backward()
nn.utils.clip_grad_norm_(parameters, 2.0)
optimizer.step()
history.append(row)
if step == 1 or step % log_every == 0 or step == steps:
_write_csv(output_dir / f"{tag}_history.csv", history)
if step % 100 == 0 or step == steps:
diagnostic_component = "band" if experiment == "D0" else "align"
print(
f"[{experiment} {tag} step {step}/{steps}] "
f"total={row['L_total']:.4f} align={row['L_align']:.4f} "
f"span={row['L_span']:.4f} band={row['L_band']:.4f} "
f"grad-{diagnostic_component}(Q/K/Z)="
f"{row.get(f'grad_{diagnostic_component}_WQ', float('nan')):.3g}/"
f"{row.get(f'grad_{diagnostic_component}_WK', float('nan')):.3g}/"
f"{row.get(f'grad_{diagnostic_component}_Z', float('nan')) if row.get(f'grad_{diagnostic_component}_Z') is not None else 'NA'}",
flush=True,
)
output_dir.mkdir(parents=True, exist_ok=True)
_write_csv(output_dir / f"{tag}_history.csv", history)
checkpoint = output_dir / f"{tag}_checkpoint.pt"
torch.save(
{
"experiment": experiment,
"method": method,
"variant": tag,
"seed": seed,
"steps": steps,
"absolute_position_encoding": absolute_position_encoding,
"model_state_dict": model.state_dict(),
"decoder_state_dict": decoder.state_dict() if decoder is not None else None,
"history": history,
},
checkpoint,
)
metric_rows, example_arrays = _evaluate_samples(
method,
experiment,
model,
eval_samples,
stats,
device=device,
batch_size=batch_size,
)
_write_csv(output_dir / f"{tag}_metrics.csv", metric_rows)
if example_arrays is not None:
np.savez_compressed(output_dir / f"{tag}_example_alignment.npz", **example_arrays)
example_sample = next(sample for sample in eval_samples if sample.sample_id == EXAMPLE_ID)
_plot_example(output_dir / tag, example_sample, example_arrays, method, experiment)
del model, decoder
if device.type == "cuda":
torch.cuda.empty_cache()
return history, metric_rows
def run(args: argparse.Namespace) -> dict[str, Any]:
started = time.time()
_seed_everything(args.seed)
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 not available")
samples = load_feature_samples(args.feature_dir, args.manifest)
if len(samples) != 100:
raise ValueError(f"debug experiments expect the complete 100-sample set, found {len(samples)}")
sample_by_id = {sample.sample_id: sample for sample in samples}
if EXAMPLE_ID not in sample_by_id:
raise ValueError(f"required debug sample is absent: {EXAMPLE_ID}")
one_sample = sample_by_id[EXAMPLE_ID]
output_root = args.output_dir
output_root.mkdir(parents=True, exist_ok=True)
full_stats = fit_feature_stats(samples)
single_stats = fit_feature_stats([one_sample])
all_metrics: list[dict[str, Any]] = []
all_history_summaries: list[dict[str, Any]] = []
specifications = [
("D0", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D0", "M4", "M4", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M3", "M3", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M4", "M4_noPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, False, False),
("D1", "M4", "M4_sinPE", [one_sample], [one_sample], single_stats, SINGLE_SAMPLE_STEPS, True, False),
("D2", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, False),
("D2", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, False),
("D3", "M3", "M3", samples, samples, full_stats, FULL_DATA_STEPS, False, True),
("D3", "M4", "M4_sinPE", samples, samples, full_stats, FULL_DATA_STEPS, True, True),
]
for experiment, method, tag, train_set, eval_set, stats, steps, use_pe, content_losses in specifications:
trial_dir = output_root / experiment
trial_dir.mkdir(parents=True, exist_ok=True)
history, metric_rows = _run_trial(
experiment,
method,
tag,
train_set,
eval_set,
stats,
trial_dir,
device=device,
seed=args.seed,
steps=steps,
absolute_position_encoding=use_pe,
batch_size=BATCH_SIZE,
include_content_losses=content_losses,
)
all_metrics.extend(metric_rows)
selected_steps = [
row
for row in history
if any(key.startswith("grad_") and value not in (None, "") for key, value in row.items())
]
summary: dict[str, Any] = {
"experiment": experiment,
"method": method,
"variant": tag,
"steps": steps,
"absolute_position_encoding": use_pe,
"final_L_total": history[-1]["L_total"],
"final_L_rec": history[-1]["L_rec"],
"final_L_con": history[-1]["L_con"],
"final_L_span": history[-1]["L_span"],
"final_L_band": history[-1]["L_band"],
"final_L_align": history[-1]["L_align"],
}
if selected_steps:
for component in ("span", "band", "align", "reconstruction", "contrastive", "total"):
for parameter in ("WQ", "WK", "Z"):
key = f"grad_{component}_{parameter}"
values = [float(row[key]) for row in selected_steps if row.get(key) not in (None, "")]
if values:
summary[f"mean_{key}"] = float(np.mean(values))
summary[f"final_{key}"] = values[-1]
summary["final_metric_rows"] = len(metric_rows)
all_history_summaries.append(summary)
print(
f"[done {experiment} {tag}] final total={summary['final_L_total']:.4f} "
f"align={summary['final_L_align']:.4f}; metrics={len(metric_rows)}",
flush=True,
)
_write_csv(output_root / "debug_summary.csv", all_history_summaries)
_write_csv(output_root / "per_sample_metrics.csv", all_metrics)
manifest = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"sample_count": len(samples),
"diagnostic_sample": EXAMPLE_ID,
"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__,
"experiments": [
{
"name": "D0",
"scope": "single sample; span + barycenter band only",
"steps_per_model": SINGLE_SAMPLE_STEPS,
"loss": "5 * L_span + 10 * L_band",
},
{
"name": "D1",
"scope": "single sample; Gaussian target KL only",
"steps_per_model": SINGLE_SAMPLE_STEPS,
"sigma_normalized_time": SIGMA,
"M4_control": "no PE versus fixed sinusoidal PE",
},
{
"name": "D2",
"scope": "all 100 clips; Gaussian target KL only; one seed; in-sample diagnostic",
"steps_per_model": FULL_DATA_STEPS,
"sigma_normalized_time": SIGMA,
"M4_position_encoding": "fixed sinusoidal",
},
{
"name": "D3",
"scope": "all 100 clips; Gaussian KL + masked reconstruction + contrastive; one seed; in-sample diagnostic",
"steps_per_model": FULL_DATA_STEPS,
"sigma_normalized_time": SIGMA,
"M4_position_encoding": "fixed sinusoidal",
},
],
"optimizer": "AdamW",
"learning_rate": LEARNING_RATE,
"dropout": 0.0,
"batch_size": BATCH_SIZE,
"gradient_logging_interval_steps": LOG_INTERVAL,
"gradient_metrics": ["W_Q", "W_K", "M4 latent slots Z"],
"features_changed": False,
"M1_M2_changed": False,
"elapsed_seconds": time.time() - started,
"interpretation_limits": [
"D0 and D1 overfit one selected sample and diagnose optimization/representability only.",
"D2 and D3 train and evaluate on the same 100 clips; they diagnose whether the target can be optimized, not generalization.",
"Gaussian targets are weak temporal priors constructed from timestamps; they are not human alignment ground truth.",
"M4 absolute sinusoidal encoding is enabled only for D1's PE control and D2/D3; prior M1-M4 and v2 results are unchanged.",
],
}
(output_root / "run_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
)
print(
f"[all done] elapsed={manifest['elapsed_seconds']:.1f}s output={output_root}",
flush=True,
)
return manifest
def build_parser() -> argparse.ArgumentParser:
project_dir = Path(__file__).resolve().parents[1]
parser = argparse.ArgumentParser(
description="Diagnose gradient flow and temporal alignment learnability for M3/M4."
)
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/alignment_debug")
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
parser.add_argument("--seed", type=int, default=42)
return parser
def main() -> int:
args = build_parser().parse_args()
run(args)
return 0
if __name__ == "__main__":
raise SystemExit(main())