Files
modeling_zhaocui/deep_learning/Q1/q1/synthetic_alignment_sanity.py
T

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()