From 54c80fb115994b1f6997544855c4903db81c4c1d Mon Sep 17 00:00:00 2001 From: wuyang <5700876+banisherwy@user.noreply.gitee.com> Date: Thu, 30 Jul 2026 10:54:08 +0800 Subject: [PATCH] research: lock AttnRes spike diagnostic runner --- experiments/k3/attnres_spike/README.md | 40 + experiments/k3/attnres_spike/scope.py | 178 ++++ experiments/k3/attnres_spike/train.py | 1098 ++++++++++++++++++++++++ 3 files changed, 1316 insertions(+) create mode 100644 experiments/k3/attnres_spike/README.md create mode 100644 experiments/k3/attnres_spike/scope.py create mode 100644 experiments/k3/attnres_spike/train.py diff --git a/experiments/k3/attnres_spike/README.md b/experiments/k3/attnres_spike/README.md new file mode 100644 index 0000000..d2920c1 --- /dev/null +++ b/experiments/k3/attnres_spike/README.md @@ -0,0 +1,40 @@ +# Attention Residuals spike-path diagnostics + +This directory implements preregistered protocol +`llm-atlas-k3-attnres-spike-path-v1`. + +It is a targeted follow-up to Round 05. It replays the exact depth-32 Block +training contract and adds diagnostic-only activation positions, gradient +reductions, and same-forward backward-rule interventions. It is not a Kimi K3 +checkpoint run and does not recover the paper's unpublished Figure 5 telemetry. + +## Frozen environment + +```text +Python /home/wuyang/.pyenv/versions/3.10.14/envs/navi-router-cu128/bin/python +PyTorch 2.11.0+cu128 +GPU NVIDIA GeForce RTX 5090 +CUBLAS_WORKSPACE_CONFIG=:4096:8 +``` + +## Step-0 smoke + +```bash +CUBLAS_WORKSPACE_CONFIG=:4096:8 \ +/home/wuyang/.pyenv/versions/3.10.14/envs/navi-router-cu128/bin/python \ + experiments/k3/attnres_spike/train.py \ + --run-kind smoke \ + --seed 2026073001 \ + --cache-dir /home/wuyang/.cache/llm-atlas/k3-attnres-gradient-scale-v1 \ + --parent-manifest experiments/k3/attnres_gradient/manifest.json \ + --manifest experiments/k3/attnres_spike/manifest.json \ + --output /home/wuyang/.cache/llm-atlas/k3-attnres-spike-path-v1/smoke/seed-2026073001.json +``` + +Formal cells use `--run-kind formal` and all three preregistered seeds. The +independent replay uses `--run-kind replay --seed 2026073001`. Formal and replay +runs are fixed to 8,000 steps; smoke performs the complete step-0 diagnostic +gate without an optimizer step. + +Raw outputs are copied into `results/raw/` only after training equivalence, +forward identity, loss-scale, reduction, and replay gates pass. diff --git a/experiments/k3/attnres_spike/scope.py b/experiments/k3/attnres_spike/scope.py new file mode 100644 index 0000000..5e20f55 --- /dev/null +++ b/experiments/k3/attnres_spike/scope.py @@ -0,0 +1,178 @@ +#!/usr/bin/env python3 +"""Rebuild the Round 06 scoping table from frozen Round 05 raw files.""" + +from __future__ import annotations + +import argparse +import json +import math +import statistics +from pathlib import Path +from typing import Any, Callable + + +SEEDS = (2026073001, 2026073002, 2026073003) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--raw-dir", + type=Path, + default=Path(__file__).resolve().parents[1] + / "attnres_gradient" + / "results" + / "raw", + ) + return parser.parse_args() + + +def pearson(left: list[float], right: list[float]) -> float: + left_mean = statistics.fmean(left) + right_mean = statistics.fmean(right) + numerator = sum( + (x - left_mean) * (y - right_mean) + for x, y in zip(left, right) + ) + left_square = sum((value - left_mean) ** 2 for value in left) + right_square = sum((value - right_mean) ** 2 for value in right) + return numerator / math.sqrt(left_square * right_square) + + +def mixer_metrics(value: dict[str, Any]) -> dict[str, float | int]: + weights = value["mean_weights"] + sources = value["sources"] + return { + "sources": sources, + "latest": weights[-1], + "maximum": max(weights), + "normalized_entropy": ( + 1.0 + if sources == 1 + else value["entropy_mean"] / math.log(sources) + ), + } + + +def average(rows: list[dict[str, Any]], key: Callable[[dict[str, Any]], float]) -> float: + return statistics.fmean(key(row) for row in rows) + + +def main() -> None: + args = parse_args() + rows = [] + source_hashes = {} + for seed in SEEDS: + path = ( + args.raw_dir + / f"formal-depth-32-block-seed-{seed}.json" + ) + run = json.loads(path.read_text()) + source_hashes[str(seed)] = run["canonical_sha256_without_self"] + diagnostic = next( + item for item in run["diagnostics"] if item["step"] == 8000 + ) + gradients = diagnostic["activation_grad_rms_by_block"] + gradient_mean = statistics.fmean(gradients) + for layer in range(32): + rows.append( + { + "seed": seed, + "layer": layer + 1, + "group": layer // 4 + 1, + "offset": layer % 4 + 1, + "normalized_gradient": gradients[layer] / gradient_mean, + "attention": mixer_metrics( + diagnostic["depth_weights"][2 * layer] + ), + "mlp": mixer_metrics( + diagnostic["depth_weights"][2 * layer + 1] + ), + } + ) + + layer_means = [] + for layer in range(1, 33): + selected = [row for row in rows if row["layer"] == layer] + layer_means.append( + { + "layer": layer, + "group": selected[0]["group"], + "offset": selected[0]["offset"], + "normalized_gradient": average( + selected, lambda row: row["normalized_gradient"] + ), + "attention_latest": average( + selected, lambda row: row["attention"]["latest"] + ), + "mlp_latest": average( + selected, lambda row: row["mlp"]["latest"] + ), + "attention_normalized_entropy": average( + selected, + lambda row: row["attention"]["normalized_entropy"], + ), + "mlp_normalized_entropy": average( + selected, lambda row: row["mlp"]["normalized_entropy"] + ), + "per_seed_normalized_gradient": [ + row["normalized_gradient"] for row in selected + ], + } + ) + + correlations = {} + for label, selected in ( + ("all_layers", rows), + ("layers_19_28", [ + row for row in rows if 19 <= row["layer"] <= 28 + ]), + ): + gradients = [row["normalized_gradient"] for row in selected] + correlations[label] = { + "points": len(selected), + "attention_latest": pearson( + gradients, + [row["attention"]["latest"] for row in selected], + ), + "mlp_latest": pearson( + gradients, [row["mlp"]["latest"] for row in selected] + ), + "attention_normalized_entropy": pearson( + gradients, + [ + row["attention"]["normalized_entropy"] + for row in selected + ], + ), + "mlp_normalized_entropy": pearson( + gradients, + [row["mlp"]["normalized_entropy"] for row in selected], + ), + "attention_maximum": pearson( + gradients, + [row["attention"]["maximum"] for row in selected], + ), + "mlp_maximum": pearson( + gradients, [row["mlp"]["maximum"] for row in selected] + ), + } + + print( + json.dumps( + { + "schema_version": 1, + "identity": "exploratory Round 05 scoping; not confirmatory Round 06", + "source_canonical_sha256": source_hashes, + "layers_19_28": layer_means[18:28], + "correlations": correlations, + }, + ensure_ascii=False, + indent=2, + sort_keys=True, + ) + ) + + +if __name__ == "__main__": + main() diff --git a/experiments/k3/attnres_spike/train.py b/experiments/k3/attnres_spike/train.py new file mode 100644 index 0000000..df8548a --- /dev/null +++ b/experiments/k3/attnres_spike/train.py @@ -0,0 +1,1098 @@ +#!/usr/bin/env python3 +"""Replay Round 05 Block training with preregistered spike-path diagnostics.""" + +from __future__ import annotations + +import argparse +import hashlib +import importlib.util +import json +import math +import os +import platform +import statistics +import sys +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Iterable + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +PROTOCOL_ID = "llm-atlas-k3-attnres-spike-path-v1" +PARENT_PROTOCOL_ID = "llm-atlas-k3-attnres-gradient-scale-v1" +DEPTH = 32 +SEEDS = (2026073001, 2026073002, 2026073003) +FORMAL_STEPS = 8000 +DIAGNOSTIC_STEPS = (0, 100, 500, 2000, 4000, 8000) +INTERVENTION_STEPS = (0, 8000) +INTERVENTION_MODES = ( + "learned", + "detached_learned", + "uniform_value_backward", +) +POSITIONS = ( + "pre_attention_input", + "attention_branch_output", + "post_attention_state", + "pre_mlp_input", + "mlp_branch_output", + "post_mlp_state", +) +REDUCTIONS = ( + "element_rms", + "token_rms_mean", + "token_rms_median", + "token_rms_p95", + "batch_mean_rms", + "token_mean_rms", + "global_l2", +) +SPIKE_LAYERS = (21, 22, 23, 24, 25) +CONTEXT = 256 +VOCABULARY = 256 +DIAGNOSTIC_WINDOWS = 16 +VALIDATION_WINDOWS = 64 +EVAL_BATCH_SIZE = 8 +TRAIN_BATCH_SIZE = 32 +TIMING_WARMUP = 20 +EPSILON = 1e-30 +SPECTRUM_TOLERANCE = 1e-6 +SCALE_TOLERANCE = 1e-5 + + +def load_parent() -> Any: + path = Path(__file__).resolve().parents[1] / "attnres_gradient" / "train.py" + spec = importlib.util.spec_from_file_location( + "k3_attnres_gradient_parent", path + ) + if spec is None or spec.loader is None: + raise RuntimeError(f"cannot import parent runner from {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +parent = load_parent() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--run-kind", choices=("smoke", "formal", "replay"), required=True) + parser.add_argument("--seed", type=int, required=True) + parser.add_argument("--cache-dir", type=Path, required=True) + parser.add_argument("--parent-manifest", type=Path, required=True) + parser.add_argument("--manifest", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + if args.seed not in SEEDS: + raise ValueError(f"seed is not preregistered: {args.seed}") + if args.run_kind == "replay" and args.seed != SEEDS[0]: + raise ValueError("the preregistered replay is seed 2026073001") + return args + + +def file_sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def canonical_sha256(value: Any) -> str: + payload = json.dumps( + value, ensure_ascii=False, sort_keys=True, separators=(",", ":") + ).encode() + return hashlib.sha256(payload).hexdigest() + + +def tensors_sha256(values: Iterable[torch.Tensor]) -> str: + digest = hashlib.sha256() + for index, value in enumerate(values): + digest.update(f"{index}\0".encode()) + digest.update(parent.tensor_bytes(value)) + return digest.hexdigest() + + +def type7_quantile(sorted_values: torch.Tensor, probability: float) -> float: + if sorted_values.ndim != 1 or sorted_values.numel() == 0: + raise ValueError("Type-7 quantile requires a non-empty flat tensor") + h = (sorted_values.numel() - 1) * probability + lower = math.floor(h) + fraction = h - lower + if fraction == 0: + value = sorted_values[lower] + else: + value = sorted_values[lower] + fraction * ( + sorted_values[lower + 1] - sorted_values[lower] + ) + return value.detach().cpu().item() + + +def mean(values: Iterable[float]) -> float: + return statistics.fmean(values) + + +def layer_statistics(values: list[float]) -> dict[str, Any]: + if len(values) != DEPTH or any(not math.isfinite(value) or value <= 0 for value in values): + raise RuntimeError("layer metric must contain 32 finite positive values") + base = parent.depth_statistics(values) + spike = [values[layer - 1] for layer in SPIKE_LAYERS] + rest = [ + value + for layer, value in enumerate(values, start=1) + if layer not in SPIKE_LAYERS + ] + ordered = sorted( + range(1, DEPTH + 1), key=lambda layer: (-values[layer - 1], layer) + ) + base.update( + { + "spike_layers": list(SPIKE_LAYERS), + "spike_mean": mean(spike), + "rest_mean": mean(rest), + "spike_contrast": mean(spike) / mean(rest), + "peak_layer": ordered[0], + "peak_normalized": values[ordered[0] - 1] / base["mean"], + "top_five_layers": ordered[:5], + "top_five_spike_overlap": len(set(ordered[:5]) & set(SPIKE_LAYERS)), + } + ) + return base + + +def gradient_reductions(gradient: torch.Tensor) -> dict[str, float]: + value = gradient.detach().float() + if value.shape != (DIAGNOSTIC_WINDOWS, CONTEXT, parent.round04.D_MODEL): + raise RuntimeError(f"unexpected gradient shape: {tuple(value.shape)}") + if not torch.isfinite(value).all(): + raise RuntimeError("non-finite activation gradient") + token_rms = value.square().mean(dim=-1).sqrt().reshape(-1) + token_rms_sorted = torch.sort(token_rms).values + result = { + "element_rms": value.square().mean().sqrt().detach().cpu().item(), + "token_rms_mean": token_rms.mean().detach().cpu().item(), + "token_rms_median": type7_quantile(token_rms_sorted, 0.5), + "token_rms_p95": type7_quantile(token_rms_sorted, 0.95), + "batch_mean_rms": value.mean(dim=0).square().mean().sqrt().detach().cpu().item(), + "token_mean_rms": value.mean(dim=1).square().mean().sqrt().detach().cpu().item(), + "global_l2": value.square().sum().sqrt().detach().cpu().item(), + } + if any(not math.isfinite(metric) or metric <= 0 for metric in result.values()): + raise RuntimeError("gradient reduction is non-finite or non-positive") + expected_l2 = result["element_rms"] * math.sqrt(value.numel()) + relative_error = abs(result["global_l2"] - expected_l2) / expected_l2 + if relative_error > 1e-6: + raise RuntimeError(f"global L2 algebraic control failed: {relative_error}") + return result + + +def rms(value: torch.Tensor) -> float: + return value.detach().float().square().mean().sqrt().cpu().item() + + +def weight_summary( + weights: torch.Tensor, + labels: list[str], + *, + mixer_index: int | None, + layer: int | None, + branch: str, + group: int | None, + offset: int | None, +) -> dict[str, Any]: + detached = weights.detach().float() + if detached.ndim != 3 or detached.shape[0] != len(labels): + raise RuntimeError("mixer weight shape/label mismatch") + source_summaries = [] + for label, source in zip(labels, detached): + flattened = torch.sort(source.reshape(-1)).values + source_summaries.append( + { + "label": label, + "mean": flattened.mean().cpu().item(), + "p05": type7_quantile(flattened, 0.05), + "median": type7_quantile(flattened, 0.5), + "p95": type7_quantile(flattened, 0.95), + } + ) + entropy = -( + detached * torch.log(detached.clamp_min(1e-30)) + ).sum(dim=0).mean().cpu().item() + source_count = len(labels) + means = [item["mean"] for item in source_summaries] + return { + "mixer_index": mixer_index, + "layer": layer, + "branch": branch, + "group": group, + "offset": offset, + "sources": source_count, + "source_summaries": source_summaries, + "entropy_mean": entropy, + "normalized_entropy": 1.0 if source_count == 1 else entropy / math.log(source_count), + "max_source_mass": max(means), + "latest_source_mass": means[-1], + } + + +def recompute_weights( + mixer: nn.Module, sources: list[torch.Tensor] +) -> torch.Tensor: + with torch.no_grad(): + values = torch.stack(sources, dim=0) + keys = mixer.key_norm(values) + logits = torch.einsum("d,nbtd->nbt", mixer.query, keys.float()) + return torch.softmax(logits, dim=0) + + +class RoutedSourceBackward(torch.autograd.Function): + @staticmethod + def forward( + ctx: Any, + values: torch.Tensor, + parent_output: torch.Tensor, + backward_weights: torch.Tensor, + ) -> torch.Tensor: + ctx.save_for_backward(values, backward_weights) + return parent_output + + @staticmethod + def backward( + ctx: Any, grad_output: torch.Tensor + ) -> tuple[torch.Tensor, None, None]: + values, backward_weights = ctx.saved_tensors + with torch.enable_grad(): + surrogate_values = values.detach().requires_grad_(True) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + surrogate_output = torch.einsum( + "nbt,nbtd->btd", + backward_weights.detach(), + surrogate_values.float(), + ).to(surrogate_values.dtype) + (grad_values,) = torch.autograd.grad( + surrogate_output, + surrogate_values, + grad_output, + retain_graph=False, + create_graph=False, + ) + return grad_values, None, None + + +@dataclass +class SpikeTrace: + positions: dict[str, list[torch.Tensor]] + mixers: list[dict[str, Any]] + + +class SpikeLanguageModel(parent.GradientLanguageModel): + """Parent model with a diagnostic-only expanded capture graph.""" + + def forward( + self, + input_ids: torch.Tensor, + capture: bool = False, + mixer_backward_mode: str = "learned", + ) -> tuple[torch.Tensor, SpikeTrace | None]: + if not capture: + if mixer_backward_mode != "learned": + raise RuntimeError("training/evaluation cannot use an intervention") + return super().forward(input_ids, capture=False) + if mixer_backward_mode not in INTERVENTION_MODES: + raise ValueError(f"unknown intervention mode: {mixer_backward_mode}") + return self._diagnostic_forward(input_ids, mixer_backward_mode) + + @staticmethod + def _capture(trace: SpikeTrace, position: str, value: torch.Tensor) -> None: + if not value.requires_grad: + raise RuntimeError(f"{position} does not require grad") + value.retain_grad() + trace.positions[position].append(value) + + @staticmethod + def _source_labels( + completed_labels: list[str], has_partial: bool + ) -> list[str]: + return completed_labels + (["current_group_partial"] if has_partial else []) + + def _mix( + self, + mixer: nn.Module, + sources: list[torch.Tensor], + labels: list[str], + *, + mode: str, + mixer_index: int | None, + layer: int | None, + branch: str, + group: int | None, + offset: int | None, + ) -> tuple[torch.Tensor, dict[str, Any]]: + parent_output, _ = mixer(sources, False) + weights = recompute_weights(mixer, sources) + summary = weight_summary( + weights, + labels, + mixer_index=mixer_index, + layer=layer, + branch=branch, + group=group, + offset=offset, + ) + if mode == "learned": + return parent_output, summary + values = torch.stack(sources, dim=0) + if mode == "detached_learned": + backward_weights = weights + else: + backward_weights = torch.full_like(weights, 1.0 / len(sources)) + routed = RoutedSourceBackward.apply( + values, parent_output, backward_weights + ) + return routed, summary + + def _diagnostic_forward( + self, input_ids: torch.Tensor, mode: str + ) -> tuple[torch.Tensor, SpikeTrace]: + trace = SpikeTrace({position: [] for position in POSITIONS}, []) + embedded = self.embed(input_ids) + completed = [embedded] + completed_labels = ["embedding"] + partial: torch.Tensor | None = None + mixer_index = 0 + + for layer_index, block in enumerate(self.blocks): + layer = layer_index + 1 + group = layer_index // 4 + 1 + offset = layer_index % 4 + 1 + + attention_sources = completed + ( + [] if partial is None else [partial] + ) + attention_labels = self._source_labels( + completed_labels, partial is not None + ) + attention_input, summary = self._mix( + self.mixers[mixer_index], + attention_sources, + attention_labels, + mode=mode, + mixer_index=mixer_index, + layer=layer, + branch="attention", + group=group, + offset=offset, + ) + trace.mixers.append(summary) + mixer_index += 1 + self._capture(trace, "pre_attention_input", attention_input) + attention_output = block.attention( + block.attention_norm(attention_input) + ) + self._capture(trace, "attention_branch_output", attention_output) + attention_for_residual = attention_output.float() + partial = ( + attention_for_residual + if partial is None + else partial + attention_for_residual + ) + self._capture(trace, "post_attention_state", partial) + + mlp_sources = completed + [partial] + mlp_labels = self._source_labels(completed_labels, True) + mlp_input, summary = self._mix( + self.mixers[mixer_index], + mlp_sources, + mlp_labels, + mode=mode, + mixer_index=mixer_index, + layer=layer, + branch="mlp", + group=group, + offset=offset, + ) + trace.mixers.append(summary) + mixer_index += 1 + self._capture(trace, "pre_mlp_input", mlp_input) + mlp_output = block.mlp(block.mlp_norm(mlp_input)) + self._capture(trace, "mlp_branch_output", mlp_output) + partial = partial + mlp_output.float() + self._capture(trace, "post_mlp_state", partial) + + if mixer_index % parent.round04.SUBLAYERS_PER_BLOCK == 0: + completed.append(partial) + completed_labels.append(f"completed_group_{group}") + partial = None + + if partial is not None or len(completed) != 9 or mixer_index != 64: + raise RuntimeError("Block aggregation topology mismatch") + if self.output_mixer is None: + raise RuntimeError("output mixer missing") + hidden, summary = self._mix( + self.output_mixer, + completed, + completed_labels, + mode=mode, + mixer_index=None, + layer=None, + branch="output", + group=None, + offset=None, + ) + trace.mixers.append(summary) + normalized = self.final_norm(hidden) + logits = F.linear(normalized, self.token_embedding.weight) + return logits, trace + + +def validate_trace(trace: SpikeTrace) -> dict[str, Any]: + expected_shape = (DIAGNOSTIC_WINDOWS, CONTEXT, parent.round04.D_MODEL) + capture = {} + for position in POSITIONS: + values = trace.positions[position] + if len(values) != DEPTH: + raise RuntimeError(f"{position} capture count mismatch") + if any(tuple(value.shape) != expected_shape for value in values): + raise RuntimeError(f"{position} capture shape mismatch") + pointers = [value.untyped_storage().data_ptr() for value in values] + if len(set(pointers)) != len(pointers): + raise RuntimeError(f"{position} contains aliased storage") + capture[position] = { + "count": len(values), + "shape": list(expected_shape), + "dtypes": [str(value.dtype) for value in values], + "storage_unique": True, + "activation_sha256": tensors_sha256(values), + "output_rms_by_layer": [rms(value) for value in values], + } + if len(trace.mixers) != 65: + raise RuntimeError("mixer capture count mismatch") + return capture + + +def run_diagnostic( + model: SpikeLanguageModel, + corpus: Any, + optimizer: torch.optim.Optimizer, + *, + mode: str, + loss_scale: float = 1.0, +) -> dict[str, Any]: + optimizer_before = parent.recursive_state_hash(optimizer.state_dict()) + model.eval() + model.zero_grad(set_to_none=True) + inputs, targets = corpus.fixed_batch( + corpus.diagnostic_starts, 0, DIAGNOSTIC_WINDOWS + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + logits, trace = model( + inputs, capture=True, mixer_backward_mode=mode + ) + unscaled_loss = parent.cross_entropy(logits, targets) + loss = unscaled_loss * loss_scale + if trace is None: + raise RuntimeError("diagnostic trace missing") + capture = validate_trace(trace) + forward = { + "logits_sha256": tensors_sha256([logits]), + "activation_sha256": { + position: capture[position]["activation_sha256"] + for position in POSITIONS + }, + "mixer_summary_sha256": canonical_sha256(trace.mixers), + "loss_nats": unscaled_loss.detach().cpu().item(), + } + loss.backward() + + positions: dict[str, Any] = {} + for position in POSITIONS: + reductions_by_layer = {name: [] for name in REDUCTIONS} + for value in trace.positions[position]: + if value.grad is None: + raise RuntimeError(f"{position} gradient is missing") + reductions = gradient_reductions(value.grad) + for name in REDUCTIONS: + reductions_by_layer[name].append(reductions[name]) + positions[position] = { + "capture": capture[position], + "reductions": { + name: { + "values": values, + "statistics": layer_statistics(values), + } + for name, values in reductions_by_layer.items() + }, + "all_gradients_present": True, + "all_gradients_finite": True, + } + + model.zero_grad(set_to_none=True) + optimizer_after = parent.recursive_state_hash(optimizer.state_dict()) + if optimizer_before != optimizer_after: + raise RuntimeError("diagnostic mutated optimizer state") + return { + "mode": mode, + "loss_scale": loss_scale, + "loss_nats": forward["loss_nats"], + "bits_per_byte": forward["loss_nats"] / math.log(2), + "forward": forward, + "positions": positions, + "mixers": trace.mixers, + "optimizer_state_before": optimizer_before, + "optimizer_state_after": optimizer_after, + "optimizer_state_unchanged": True, + } + + +def forward_identity_gate(modes: dict[str, Any]) -> dict[str, Any]: + learned = modes["learned"]["forward"] + comparisons = {} + for name in INTERVENTION_MODES[1:]: + other = modes[name]["forward"] + comparisons[name] = { + "logits_exact": other["logits_sha256"] == learned["logits_sha256"], + "loss_exact": other["loss_nats"] == learned["loss_nats"], + "activations_exact": ( + other["activation_sha256"] == learned["activation_sha256"] + ), + "mixer_summaries_exact": ( + other["mixer_summary_sha256"] + == learned["mixer_summary_sha256"] + ), + } + passed = all(all(values.values()) for values in comparisons.values()) + if not passed: + raise RuntimeError(f"forward identity gate failed: {comparisons}") + return {"passed": True, "comparisons": comparisons} + + +def initialization_negative_control(modes: dict[str, Any]) -> dict[str, Any]: + learned = modes["learned"] + comparisons = {} + passed = True + for mode in INTERVENTION_MODES[1:]: + per_position = {} + for position in POSITIONS: + left = learned["positions"][position]["reductions"]["element_rms"] + right = modes[mode]["positions"][position]["reductions"]["element_rms"] + raw_errors = [ + abs(a - b) / a + for a, b in zip(left["values"], right["values"]) + ] + normalized_errors = [ + abs(a - b) + for a, b in zip( + left["statistics"]["normalized"], + right["statistics"]["normalized"], + ) + ] + position_passed = ( + all(value > 0 and math.isfinite(value) for value in left["values"]) + and max(raw_errors) <= SPECTRUM_TOLERANCE + and max(normalized_errors) <= SPECTRUM_TOLERANCE + ) + passed = passed and position_passed + per_position[position] = { + "passed": position_passed, + "max_raw_relative_error": max(raw_errors), + "max_normalized_absolute_error": max(normalized_errors), + } + comparisons[mode] = per_position + if not passed: + raise RuntimeError( + "initialization negative control failed: " + + json.dumps(comparisons, sort_keys=True) + ) + return {"passed": True, "comparisons": comparisons} + + +def loss_scale_gate(base: dict[str, Any], doubled: dict[str, Any]) -> dict[str, Any]: + if base["forward"] != doubled["forward"]: + raise RuntimeError("loss scaling unexpectedly changed the forward pass") + checks = {} + passed = True + for position in POSITIONS: + checks[position] = {} + for reduction in REDUCTIONS: + left = base["positions"][position]["reductions"][reduction] + right = doubled["positions"][position]["reductions"][reduction] + scale_errors = [ + abs((b / a) - 2.0) + for a, b in zip(left["values"], right["values"]) + ] + normalized_errors = [ + abs(a - b) + for a, b in zip( + left["statistics"]["normalized"], + right["statistics"]["normalized"], + ) + ] + cv_error = abs( + left["statistics"]["population_cv"] + - right["statistics"]["population_cv"] + ) + contrast_error = abs( + left["statistics"]["spike_contrast"] + - right["statistics"]["spike_contrast"] + ) + item_passed = ( + max(scale_errors) <= SCALE_TOLERANCE + and max(normalized_errors) <= SPECTRUM_TOLERANCE + and cv_error <= SPECTRUM_TOLERANCE + and contrast_error <= SPECTRUM_TOLERANCE + ) + passed = passed and item_passed + checks[position][reduction] = { + "passed": item_passed, + "max_scale_ratio_error": max(scale_errors), + "max_normalized_absolute_error": max(normalized_errors), + "cv_absolute_error": cv_error, + "spike_contrast_absolute_error": contrast_error, + } + if not passed: + raise RuntimeError("loss-scale gate failed") + return {"passed": True, "checks": checks} + + +def run_diagnostic_bundle( + model: SpikeLanguageModel, + corpus: Any, + optimizer: torch.optim.Optimizer, + step: int, +) -> dict[str, Any]: + learned = run_diagnostic( + model, corpus, optimizer, mode="learned" + ) + modes = {"learned": learned} + if step in INTERVENTION_STEPS: + for mode in INTERVENTION_MODES[1:]: + modes[mode] = run_diagnostic( + model, corpus, optimizer, mode=mode + ) + result = {"step": step, "modes": modes} + if step in INTERVENTION_STEPS: + result["forward_identity_gate"] = forward_identity_gate(modes) + if step == 0: + result["initialization_negative_control"] = ( + initialization_negative_control(modes) + ) + doubled = run_diagnostic( + model, corpus, optimizer, mode="learned", loss_scale=2.0 + ) + result["loss_scale_gate"] = loss_scale_gate(learned, doubled) + model.zero_grad(set_to_none=True) + return result + + +def frozen_training_compare( + result: dict[str, Any], + parent_raw: dict[str, Any], +) -> dict[str, Any]: + checks = { + "final_model_state": ( + result["hashes"]["final_model_state"] + == parent_raw["hashes"]["final_model_state"] + ), + "final_optimizer_state": ( + result["hashes"]["final_optimizer_state"] + == parent_raw["hashes"]["final_optimizer_state"] + ), + "evaluations": result["evaluations"] == parent_raw["evaluations"], + "training_history": ( + result["training_history"] == parent_raw["training_history"] + ), + } + post_mlp_exact = [] + for new, old in zip(result["diagnostics"], parent_raw["diagnostics"]): + new_values = new["modes"]["learned"]["positions"]["post_mlp_state"][ + "reductions" + ]["element_rms"]["values"] + post_mlp_exact.append( + new["step"] == old["step"] + and new_values == old["activation_grad_rms_by_block"] + ) + checks["post_mlp_element_rms_all_steps"] = all(post_mlp_exact) + passed = all(checks.values()) + if not passed: + raise RuntimeError(f"Round 05 training equivalence failed: {checks}") + return {"passed": True, "checks": checks} + + +def load_and_verify_inputs( + args: argparse.Namespace, +) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any], Path]: + manifest = json.loads(args.manifest.read_text()) + parent_manifest = json.loads(args.parent_manifest.read_text()) + repo_root = Path(__file__).resolve().parents[3] + if manifest["protocol_id"] != PROTOCOL_ID: + raise RuntimeError("Round 06 manifest protocol mismatch") + if parent_manifest["protocol_id"] != PARENT_PROTOCOL_ID: + raise RuntimeError("parent manifest protocol mismatch") + frozen_fields = { + "architecture": "block", + "depth": DEPTH, + "formal_seeds": list(SEEDS), + "positions": list(POSITIONS), + } + for key, expected_value in frozen_fields.items(): + if manifest[key] != expected_value: + raise RuntimeError(f"Round 06 manifest field mismatch: {key}") + if manifest["training"]["steps"] != FORMAL_STEPS: + raise RuntimeError("Round 06 formal-step mismatch") + if manifest["training"]["diagnostic_steps"] != list(DIAGNOSTIC_STEPS): + raise RuntimeError("Round 06 diagnostic-step mismatch") + if manifest["interventions"]["modes"] != list(INTERVENTION_MODES): + raise RuntimeError("Round 06 intervention-mode mismatch") + if manifest["interventions"]["steps"] != list(INTERVENTION_STEPS): + raise RuntimeError("Round 06 intervention-step mismatch") + manifest_reductions = ( + manifest["reductions"]["confirmatory"] + + manifest["reductions"]["exploratory"] + + manifest["reductions"]["algebraic_control"] + ) + if manifest_reductions != list(REDUCTIONS): + raise RuntimeError("Round 06 reduction list mismatch") + expected_parent_hash = manifest["parent_artifacts"]["manifest_sha256"] + if file_sha256(args.parent_manifest) != expected_parent_hash: + raise RuntimeError("parent manifest file hash mismatch") + if file_sha256(Path(parent.__file__)) != manifest["parent_artifacts"]["runner_sha256"]: + raise RuntimeError("parent runner file hash mismatch") + parent_protocol_path = ( + repo_root / "research" / "K3_ATTNRES_GRADIENT_SCALE_PROTOCOL.md" + ) + if ( + file_sha256(parent_protocol_path) + != manifest["parent_artifacts"]["protocol_sha256"] + ): + raise RuntimeError("parent protocol file hash mismatch") + for key in ( + "formal_schedule_sha256", + "validation_tensor_sha256", + "diagnostic_tensor_sha256", + ): + if parent_manifest["windows"][key] != manifest["parent_artifacts"][key]: + raise RuntimeError(f"parent data hash mismatch: {key}") + parent_raw_path = ( + repo_root + / "experiments" + / "k3" + / "attnres_gradient" + / "results" + / "raw" + / f"formal-depth-32-block-seed-{args.seed}.json" + ) + expected = manifest["round05_expected"][str(args.seed)] + if file_sha256(parent_raw_path) != expected["raw_file_sha256"]: + raise RuntimeError("Round 05 raw physical hash mismatch") + parent_raw = json.loads(parent_raw_path.read_text()) + if parent_raw["canonical_sha256_without_self"] != expected["canonical_sha256"]: + raise RuntimeError("Round 05 raw canonical hash mismatch") + if parent_raw["hashes"]["final_model_state"] != expected["final_model_state"]: + raise RuntimeError("Round 05 expected model-state hash mismatch") + if ( + parent_raw["hashes"]["final_optimizer_state"] + != expected["final_optimizer_state"] + ): + raise RuntimeError("Round 05 expected optimizer-state hash mismatch") + return manifest, parent_manifest, parent_raw, repo_root + + +def main() -> None: + args = parse_args() + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required") + if os.environ.get("CUBLAS_WORKSPACE_CONFIG") != ":4096:8": + raise RuntimeError("CUBLAS_WORKSPACE_CONFIG must be :4096:8") + manifest, parent_manifest, parent_raw, repo_root = load_and_verify_inputs(args) + parent.configure_round04_globals(DEPTH) + parent.configure_determinism(args.seed) + corpus = parent.round04.ByteCorpus( + args.cache_dir, parent_manifest, torch.device("cuda") + ) + model = SpikeLanguageModel("block").to(torch.device("cuda")) + + initial_public_hash = parent.named_state_hash( + model, include_mixers=False + ) + initial_mixer_hash = parent.named_state_hash( + model, include_mixers=True + ) + public_structure_hash, public_tensors, public_elements = ( + parent.state_structure_hash(model, include_mixers=False) + ) + input_gate_hashes = parent.model_input_gate_hashes( + corpus, parent_manifest, args.seed, TRAIN_BATCH_SIZE + ) + + decay_parameters: list[nn.Parameter] = [] + no_decay_parameters: list[nn.Parameter] = [] + for parameter in model.parameters(): + (decay_parameters if parameter.ndim >= 2 else no_decay_parameters).append( + parameter + ) + optimizer = torch.optim.AdamW( + [ + {"params": decay_parameters, "weight_decay": parent.WEIGHT_DECAY}, + {"params": no_decay_parameters, "weight_decay": 0.0}, + ], + lr=parent.PEAK_LR, + betas=parent.BETAS, + eps=parent.ADAM_EPS, + ) + + evaluations = [ + { + "step": 0, + **parent.evaluate( + model, corpus, VALIDATION_WINDOWS, EVAL_BATCH_SIZE + ), + } + ] + diagnostics = [run_diagnostic_bundle(model, corpus, optimizer, 0)] + print( + json.dumps( + { + "event": "diagnostic", + "step": 0, + "seed": args.seed, + "post_mlp_cv": diagnostics[0]["modes"]["learned"][ + "positions" + ]["post_mlp_state"]["reductions"]["element_rms"][ + "statistics" + ]["population_cv"], + }, + sort_keys=True, + ), + flush=True, + ) + + steps = 0 if args.run_kind == "smoke" else FORMAL_STEPS + training_history: list[dict[str, float | int]] = [] + step_times: list[float] = [] + if steps: + model.train() + for step in range(1, steps + 1): + lr = parent.learning_rate(step, steps) + for group in optimizer.param_groups: + group["lr"] = lr + inputs, targets = corpus.training_batch( + args.seed, step, TRAIN_BATCH_SIZE + ) + optimizer.zero_grad(set_to_none=True) + torch.cuda.synchronize() + started = time.perf_counter() + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + logits, trace = model(inputs) + if trace is not None: + raise RuntimeError("training unexpectedly captured a trace") + loss = parent.cross_entropy(logits, targets) + if not torch.isfinite(loss): + raise RuntimeError(f"non-finite loss at step {step}") + loss.backward() + unclipped_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), parent.GRAD_CLIP + ) + optimizer.step() + torch.cuda.synchronize() + elapsed_ms = (time.perf_counter() - started) * 1000 + if step == TIMING_WARMUP: + torch.cuda.reset_peak_memory_stats() + elif step > TIMING_WARMUP: + step_times.append(elapsed_ms) + if step == 1 or step % 10 == 0 or step == steps: + training_history.append( + { + "step": step, + "loss_nats": loss.detach().cpu().item(), + "bits_per_byte": loss.detach().cpu().item() / math.log(2), + "learning_rate": lr, + "unclipped_grad_norm": float( + unclipped_norm.detach().cpu() + ), + } + ) + if step in DIAGNOSTIC_STEPS: + evaluations.append( + { + "step": step, + **parent.evaluate( + model, + corpus, + VALIDATION_WINDOWS, + EVAL_BATCH_SIZE, + ), + } + ) + diagnostics.append( + run_diagnostic_bundle(model, corpus, optimizer, step) + ) + post = diagnostics[-1]["modes"]["learned"]["positions"][ + "post_mlp_state" + ]["reductions"]["element_rms"]["statistics"] + print( + json.dumps( + { + "event": "diagnostic", + "step": step, + "seed": args.seed, + "validation_bpc": evaluations[-1][ + "bits_per_byte" + ], + "post_mlp_cv": post["population_cv"], + "spike_contrast": post["spike_contrast"], + }, + sort_keys=True, + ), + flush=True, + ) + model.train() + + timing = { + "warmup_steps_excluded": TIMING_WARMUP, + "measured_steps": len(step_times), + "mean_ms": mean(step_times) if step_times else None, + "median_ms": statistics.median(step_times) if step_times else None, + "p95_ms": ( + type7_quantile(torch.tensor(sorted(step_times)), 0.95) + if step_times + else None + ), + "peak_allocated_bytes": torch.cuda.max_memory_allocated(), + "peak_reserved_bytes": torch.cuda.max_memory_reserved(), + } + result = { + "schema_version": 1, + "protocol_id": PROTOCOL_ID, + "parent_protocol_id": PARENT_PROTOCOL_ID, + "run_kind": args.run_kind, + "architecture": "block", + "depth": DEPTH, + "seed": args.seed, + "steps": steps, + "batch_size": TRAIN_BATCH_SIZE, + "target_bytes_seen": steps * TRAIN_BATCH_SIZE * CONTEXT, + "manifest": { + "path": str(args.manifest), + "file_sha256": file_sha256(args.manifest), + "parent_path": str(args.parent_manifest), + "parent_file_sha256": file_sha256(args.parent_manifest), + "formal_schedule_sha256": parent_manifest["windows"][ + "formal_schedule_sha256" + ], + "validation_tensor_sha256": parent_manifest["windows"][ + "validation_tensor_sha256" + ], + "diagnostic_tensor_sha256": parent_manifest["windows"][ + "diagnostic_tensor_sha256" + ], + "input_gate_tensor_hashes": input_gate_hashes, + }, + "model": { + "layers": DEPTH, + "aggregation_groups": 8, + "blocks_per_group": 4, + "d_model": parent.round04.D_MODEL, + "heads": parent.round04.HEADS, + "d_ff": parent.round04.D_FF, + "parameters": parent.parameter_inventory(model), + }, + "optimizer": { + "name": "AdamW", + "betas": list(parent.BETAS), + "epsilon": parent.ADAM_EPS, + "weight_decay_ndim_ge_2": parent.WEIGHT_DECAY, + "peak_lr": parent.PEAK_LR, + "min_lr": parent.MIN_LR, + "warmup_steps": parent.WARMUP_STEPS, + "grad_clip": parent.GRAD_CLIP, + }, + "hashes": { + "initial_public_parameter_structure": public_structure_hash, + "initial_public_parameter_tensors": public_tensors, + "initial_public_parameter_elements": public_elements, + "initial_public_parameters": initial_public_hash, + "initial_mixer_parameters": initial_mixer_hash, + "final_public_parameters": parent.named_state_hash( + model, include_mixers=False + ), + "final_mixer_parameters": parent.named_state_hash( + model, include_mixers=True + ), + "final_model_state": parent.named_state_hash( + model, include_mixers=None + ), + "final_optimizer_state": parent.recursive_state_hash( + optimizer.state_dict() + ), + }, + "evaluations": evaluations, + "diagnostics": diagnostics, + "training_history": training_history, + "timing": timing, + "environment": { + "python": platform.python_version(), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "gpu": torch.cuda.get_device_name(0), + "compute_capability": list(torch.cuda.get_device_capability(0)), + "cublas_workspace_config": os.environ[ + "CUBLAS_WORKSPACE_CONFIG" + ], + "deterministic_algorithms": ( + torch.are_deterministic_algorithms_enabled() + ), + "autocast": "cuda-bfloat16-forward-fp32-cross-entropy", + "compile": False, + }, + "artifacts": { + "runner_sha256": file_sha256(Path(__file__)), + "protocol_sha256": file_sha256( + repo_root / "research" / "K3_ATTNRES_SPIKE_PROTOCOL.md" + ), + "scoping_sha256": file_sha256( + repo_root / "research" / "K3_ATTNRES_SPIKE_SCOPING.md" + ), + }, + } + if steps: + result["round05_equivalence"] = frozen_training_compare( + result, parent_raw + ) + else: + result["round05_equivalence"] = None + result["canonical_sha256_without_self"] = canonical_sha256(result) + args.output.parent.mkdir(parents=True, exist_ok=True) + temporary = args.output.with_suffix(args.output.suffix + ".tmp") + temporary.write_text( + json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True) + + "\n" + ) + os.replace(temporary, args.output) + print( + json.dumps( + { + "output": str(args.output), + "run_kind": args.run_kind, + "seed": args.seed, + "steps": steps, + "final_bpc": evaluations[-1]["bits_per_byte"], + "canonical_sha256": result[ + "canonical_sha256_without_self" + ], + "timing": timing, + }, + ensure_ascii=False, + indent=2, + ), + flush=True, + ) + + +if __name__ == "__main__": + main()