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