339 lines
13 KiB
Python
339 lines
13 KiB
Python
"""Fit Gaussian time bands with synthetic time-only inputs as an attention control."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import csv
|
|
import json
|
|
import platform
|
|
import random
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import matplotlib
|
|
|
|
matplotlib.use("Agg")
|
|
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
import torch
|
|
from torch import Tensor, nn
|
|
|
|
from .metrics import alignment_trajectory, normalized_attention_entropy
|
|
from .models import SharedLatentTimeline, TextAnchoredCrossAttention
|
|
from .types import MODALITIES, SequenceBatch
|
|
|
|
|
|
GRID_SIZE = 50
|
|
HIDDEN_SIZE = 128
|
|
SIGMA = 0.10
|
|
SEED = 42
|
|
MODALITY_LENGTHS = {"text": GRID_SIZE, "audio": 256, "vision": 100}
|
|
|
|
|
|
def _seed(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 _synthetic_sequence(length: int, device: torch.device) -> SequenceBatch:
|
|
times = torch.linspace(0.0, 1.0, length, device=device).unsqueeze(0)
|
|
# These features contain only the source's synthetic time coordinate.
|
|
features = torch.stack(
|
|
(times, times.square(), torch.ones_like(times)), dim=-1
|
|
)
|
|
valid = torch.ones_like(times, dtype=torch.bool)
|
|
return SequenceBatch(features=features, times=times, valid=valid)
|
|
|
|
|
|
def _targets(
|
|
sequences: dict[str, SequenceBatch], method: str, device: torch.device
|
|
) -> dict[str, Tensor]:
|
|
centers = (torch.arange(GRID_SIZE, device=device, dtype=torch.float32) + 0.5) / GRID_SIZE
|
|
names = ("audio", "vision") if method == "M3" else MODALITIES
|
|
targets: dict[str, Tensor] = {}
|
|
for name in names:
|
|
times = sequences[name].times
|
|
logits = -0.5 * ((times[:, None, :] - centers[None, :, None]) / SIGMA).square()
|
|
targets[name] = torch.softmax(logits, dim=-1)
|
|
return targets
|
|
|
|
|
|
def _loss(output: Any, targets: dict[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)
|
|
losses.append((safe_target * (safe_target.log() - predicted.log())).sum(-1).mean())
|
|
return torch.stack(losses).mean()
|
|
|
|
|
|
def _attention_layers(model: nn.Module, method: str) -> list[nn.MultiheadAttention]:
|
|
if method == "M3":
|
|
return [model.audio_attention.attention, model.vision_attention.attention]
|
|
return [model.attention[name].attention for name in MODALITIES]
|
|
|
|
|
|
def _gradient_summary(
|
|
model: nn.Module, method: str, loss: Tensor
|
|
) -> dict[str, float]:
|
|
layers = _attention_layers(model, method)
|
|
params = [layer.in_proj_weight for layer in layers]
|
|
if method == "M4":
|
|
params.append(model.slots)
|
|
gradients = torch.autograd.grad(loss, params, retain_graph=True)
|
|
q_sq = torch.zeros((), device=loss.device)
|
|
k_sq = torch.zeros((), device=loss.device)
|
|
for grad in gradients[: len(layers)]:
|
|
q_sq += grad[:HIDDEN_SIZE].square().sum()
|
|
k_sq += grad[HIDDEN_SIZE : 2 * HIDDEN_SIZE].square().sum()
|
|
result = {"grad_WQ": float(q_sq.sqrt().item()), "grad_WK": float(k_sq.sqrt().item())}
|
|
if method == "M4":
|
|
result["grad_Z"] = float(gradients[-1].norm().item())
|
|
return result
|
|
|
|
|
|
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
|
if not rows:
|
|
return
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
with path.open("w", newline="", encoding="utf-8-sig") as handle:
|
|
fields = list(dict.fromkeys(key for row in rows for key in row))
|
|
writer = csv.DictWriter(handle, fieldnames=fields)
|
|
writer.writeheader()
|
|
writer.writerows(rows)
|
|
|
|
|
|
def _save_plots(
|
|
folder: Path,
|
|
method: str,
|
|
variant: str,
|
|
output: Any,
|
|
targets: dict[str, Tensor],
|
|
sequences: dict[str, SequenceBatch],
|
|
duration: float = 1.0,
|
|
) -> list[dict[str, Any]]:
|
|
names = ("audio", "vision") if method == "M3" else MODALITIES
|
|
rows: list[dict[str, Any]] = []
|
|
fig, axes = plt.subplots(len(names), 2, figsize=(11, 3.4 * len(names)), constrained_layout=True)
|
|
if len(names) == 1:
|
|
axes = np.asarray([axes])
|
|
fig_traj, ax_traj = plt.subplots(figsize=(8, 5), constrained_layout=True)
|
|
grid = (np.arange(GRID_SIZE, dtype=np.float32) + 0.5) / GRID_SIZE
|
|
for row_index, name in enumerate(names):
|
|
weights = output.weights[name][0].detach().cpu().numpy()
|
|
target = targets[name][0].detach().cpu().numpy()
|
|
times = sequences[name].times[0].detach().cpu().numpy()
|
|
entropy = normalized_attention_entropy(
|
|
output.weights[name], sequences[name].valid
|
|
).mean().item()
|
|
trajectory = alignment_trajectory(
|
|
output.weights[name], sequences[name].times, torch.tensor([duration], device=output.weights[name].device)
|
|
)[0]
|
|
trajectory_np = trajectory.detach().cpu().numpy()
|
|
span = float(trajectory_np[-1] - trajectory_np[0])
|
|
kl = float(
|
|
(targets[name] * (targets[name].clamp_min(1e-12).log() - output.weights[name].clamp_min(1e-8).log()))
|
|
.sum(-1)
|
|
.mean()
|
|
.item()
|
|
)
|
|
rows.append(
|
|
{
|
|
"method": method,
|
|
"variant": variant,
|
|
"modality": name,
|
|
"normalized_entropy": entropy,
|
|
"trajectory_span": span,
|
|
"mean_absolute_time_center_error": float(
|
|
np.mean(np.abs(trajectory_np - grid))
|
|
),
|
|
"gaussian_target_kl": kl,
|
|
}
|
|
)
|
|
extent = (float(times[0]), float(times[-1]), 0.0, 1.0)
|
|
image = axes[row_index, 0].imshow(
|
|
weights, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma"
|
|
)
|
|
axes[row_index, 0].set_title(f"{name}: learned A")
|
|
axes[row_index, 0].set_xlabel("synthetic source time")
|
|
axes[row_index, 0].set_ylabel("slot / K")
|
|
fig.colorbar(image, ax=axes[row_index, 0], fraction=0.046, pad=0.04)
|
|
image_target = axes[row_index, 1].imshow(
|
|
target, origin="lower", aspect="auto", interpolation="nearest", extent=extent, cmap="magma"
|
|
)
|
|
axes[row_index, 1].set_title(f"{name}: Gaussian target P")
|
|
axes[row_index, 1].set_xlabel("synthetic source time")
|
|
axes[row_index, 1].set_ylabel("slot / K")
|
|
fig.colorbar(image_target, ax=axes[row_index, 1], fraction=0.046, pad=0.04)
|
|
ax_traj.plot(grid, trajectory_np, label=f"{name} learned")
|
|
ax_traj.plot([0, 1], [0, 1], "k:", label="uniform-time reference")
|
|
ax_traj.set(xlim=(0, 1), ylim=(0, 1), xlabel="slot position", ylabel="expected source time")
|
|
ax_traj.grid(alpha=0.2)
|
|
ax_traj.legend()
|
|
ax_traj.set_title(f"{method} {variant}: synthetic alignment trajectory")
|
|
folder.mkdir(parents=True, exist_ok=True)
|
|
fig.savefig(folder / f"{variant}_heatmap.png", dpi=160)
|
|
fig_traj.savefig(folder / f"{variant}_trajectory.png", dpi=160)
|
|
plt.close(fig)
|
|
plt.close(fig_traj)
|
|
return rows
|
|
|
|
|
|
def _run_trial(
|
|
method: str,
|
|
variant: str,
|
|
*,
|
|
device: torch.device,
|
|
steps: int,
|
|
seed: int,
|
|
output_dir: Path,
|
|
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]:
|
|
_seed(seed)
|
|
dimensions = {name: 3 for name in MODALITIES}
|
|
if method == "M3":
|
|
model: nn.Module = TextAnchoredCrossAttention(
|
|
dimensions, grid_size=GRID_SIZE, hidden_size=HIDDEN_SIZE, heads=4, dropout=0.0
|
|
)
|
|
use_pe = False
|
|
else:
|
|
use_pe = variant == "M4_sinPE"
|
|
model = SharedLatentTimeline(
|
|
dimensions,
|
|
grid_size=GRID_SIZE,
|
|
hidden_size=HIDDEN_SIZE,
|
|
heads=4,
|
|
dropout=0.0,
|
|
absolute_position_encoding=use_pe,
|
|
)
|
|
model.to(device)
|
|
sequences = {
|
|
name: _synthetic_sequence(MODALITY_LENGTHS[name], device) for name in MODALITIES
|
|
}
|
|
targets = _targets(sequences, method, device)
|
|
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.0)
|
|
history: list[dict[str, Any]] = []
|
|
output = None
|
|
print(f"[{variant}] synthetic time-only inputs, steps={steps}, device={device}", flush=True)
|
|
for step in range(1, steps + 1):
|
|
model.train()
|
|
output = model(sequences)
|
|
loss = _loss(output, targets)
|
|
if not torch.isfinite(loss):
|
|
raise FloatingPointError(f"non-finite synthetic alignment loss for {variant} at step {step}")
|
|
row: dict[str, Any] = {"method": method, "variant": variant, "step": step, "L_align": float(loss.item())}
|
|
if step == 1 or step % 100 == 0 or step == steps:
|
|
row.update(_gradient_summary(model, method, loss))
|
|
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:
|
|
print(
|
|
f"[{variant} {step}/{steps}] KL={row['L_align']:.5f} "
|
|
f"grad_Q/K/Z={row.get('grad_WQ', 0):.3g}/{row.get('grad_WK', 0):.3g}/"
|
|
f"{row.get('grad_Z', float('nan')):.3g}",
|
|
flush=True,
|
|
)
|
|
assert output is not None
|
|
model.eval()
|
|
with torch.no_grad():
|
|
output = model(sequences)
|
|
metric_rows = _save_plots(
|
|
output_dir / variant, method, variant, output, targets, sequences
|
|
)
|
|
history_rows = history
|
|
checkpoint = output_dir / f"{variant}_checkpoint.pt"
|
|
torch.save(
|
|
{
|
|
"method": method,
|
|
"variant": variant,
|
|
"seed": seed,
|
|
"steps": steps,
|
|
"absolute_position_encoding": use_pe,
|
|
"model_state_dict": model.state_dict(),
|
|
"synthetic_only": True,
|
|
},
|
|
checkpoint,
|
|
)
|
|
return metric_rows, history_rows, {"checkpoint": str(checkpoint), "final_loss": history[-1]["L_align"]}
|
|
|
|
|
|
def run(args: argparse.Namespace) -> dict[str, Any]:
|
|
output_dir = args.output_dir
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
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")
|
|
variants = [("M3", "M3"), ("M4", "M4_noPE"), ("M4", "M4_sinPE")]
|
|
all_metrics: list[dict[str, Any]] = []
|
|
all_history: list[dict[str, Any]] = []
|
|
model_summaries: list[dict[str, Any]] = []
|
|
for method, variant in variants:
|
|
metrics, history, summary = _run_trial(
|
|
method, variant, device=device, steps=args.steps, seed=args.seed, output_dir=output_dir
|
|
)
|
|
all_metrics.extend(metrics)
|
|
all_history.extend(history)
|
|
model_summaries.append({"method": method, "variant": variant, **summary})
|
|
_write_csv(output_dir / "metrics.csv", all_metrics)
|
|
_write_csv(output_dir / "training_history.csv", all_history)
|
|
manifest = {
|
|
"created_utc": datetime.now(timezone.utc).isoformat(),
|
|
"purpose": "Check whether the existing M3/M4 attention implementation can fit a time band when all inputs encode only synthetic time.",
|
|
"sample_count": 1,
|
|
"synthetic_only": True,
|
|
"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,
|
|
"feature_rule": "[t, t^2, 1] per source position; no extracted text/audio/video content is used",
|
|
"sequence_lengths": MODALITY_LENGTHS,
|
|
"optimizer": "AdamW",
|
|
"learning_rate": 1e-3,
|
|
"steps_per_variant": args.steps,
|
|
"variants": model_summaries,
|
|
"interpretation_limits": [
|
|
"This is an implementation/optimization control only; it does not evaluate real features or alignment accuracy.",
|
|
"The time-only features deliberately provide source-position information that the real feature inputs may not contain.",
|
|
],
|
|
}
|
|
(output_dir / "run_manifest.json").write_text(
|
|
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
print(f"[all done] wrote synthetic control to {output_dir}", 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=SEED)
|
|
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
|
parser.add_argument(
|
|
"--output-dir", type=Path, default=project / "outputs/alignment_debug/synthetic"
|
|
)
|
|
return parser
|
|
|
|
|
|
def main() -> None:
|
|
args = build_parser().parse_args()
|
|
run(args)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|