777 lines
30 KiB
Python
777 lines
30 KiB
Python
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())
|