From f99872039f08bd72f845bcdc29c1dbbc686fe694 Mon Sep 17 00:00:00 2001 From: wuyang <5700876+banisherwy@user.noreply.gitee.com> Date: Thu, 30 Jul 2026 06:38:26 +0800 Subject: [PATCH] research: add reduced AttnRes runner --- experiments/k3/attnres/README.md | 43 ++ experiments/k3/attnres/build_dataset.py | 192 +++++++ experiments/k3/attnres/train.py | 715 ++++++++++++++++++++++++ 3 files changed, 950 insertions(+) create mode 100644 experiments/k3/attnres/README.md create mode 100644 experiments/k3/attnres/build_dataset.py create mode 100644 experiments/k3/attnres/train.py diff --git a/experiments/k3/attnres/README.md b/experiments/k3/attnres/README.md new file mode 100644 index 0000000..66b1614 --- /dev/null +++ b/experiments/k3/attnres/README.md @@ -0,0 +1,43 @@ +# Reduced Attention Residuals reproduction + +This directory implements protocol +`llm-atlas-k3-attnres-reduced-v1`, frozen in +`research/K3_ATTNRES_REDUCED_PROTOCOL.md`. + +The experiment is a reduced independent mechanism probe. It is not a Kimi K3 +checkpoint forward pass and not a reproduction of the paper-scale training run. + +## Environment + +The pinned execution environment used by this project is: + +```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 +``` + +## Build the frozen dataset + +```bash +python experiments/k3/attnres/build_dataset.py \ + --cache-dir /home/wuyang/.cache/llm-atlas/k3-attnres-reduced-v1 \ + --manifest experiments/k3/attnres/manifest.json +``` + +## Run one cell + +```bash +CUBLAS_WORKSPACE_CONFIG=:4096:8 \ +python experiments/k3/attnres/train.py \ + --architecture baseline \ + --seed 2026073001 \ + --cache-dir /home/wuyang/.cache/llm-atlas/k3-attnres-reduced-v1 \ + --manifest experiments/k3/attnres/manifest.json \ + --output /home/wuyang/.cache/llm-atlas/k3-attnres-reduced-v1/runs/baseline-2026073001.json +``` + +Raw parquet and checkpoints stay in the local cache. Frozen manifests, metric +JSON, analyses, code, checksums, and a compact website payload enter the public +repository. diff --git a/experiments/k3/attnres/build_dataset.py b/experiments/k3/attnres/build_dataset.py new file mode 100644 index 0000000..7cdcff7 --- /dev/null +++ b/experiments/k3/attnres/build_dataset.py @@ -0,0 +1,192 @@ +#!/usr/bin/env python3 +"""Download and freeze the byte-level WikiText-2 corpus for the AttnRes study.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import urllib.request +from pathlib import Path +from typing import Any + +import pyarrow.parquet as pq + + +PROTOCOL_ID = "llm-atlas-k3-attnres-reduced-v1" +DATASET_REPO = "Salesforce/wikitext" +DATASET_REVISION = "b08601e04326c79dfdd32d625aee71d232d685c3" +DATASET_VARIANT = "wikitext-2-raw-v1" +SPLITS = ("train", "validation", "test") +SEEDS = (2026073001, 2026073002, 2026073003) +CONTEXT = 256 +FORMAL_STEPS = 2000 +FORMAL_BATCH = 32 +VALIDATION_WINDOWS = 64 +DIAGNOSTIC_WINDOWS = 16 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--cache-dir", type=Path, required=True) + parser.add_argument("--manifest", type=Path, required=True) + return parser.parse_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 bytes_sha256(payload: bytes) -> str: + return hashlib.sha256(payload).hexdigest() + + +def atomic_json(path: Path, value: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text( + json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n" + ) + os.replace(temporary, path) + + +def download(url: str, path: Path) -> None: + if path.exists(): + return + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".part") + request = urllib.request.Request( + url, + headers={"User-Agent": "llm-atlas-k3-attnres-reduced/1.0"}, + ) + with urllib.request.urlopen(request, timeout=120) as response: + with temporary.open("wb") as output: + while block := response.read(1024 * 1024): + output.write(block) + os.replace(temporary, path) + + +def window_start(label: str, index: int, corpus_length: int, seed: int | None = None) -> int: + fields = [PROTOCOL_ID, label] + if seed is not None: + fields.append(str(seed)) + fields.append(str(index)) + payload = "\0".join(fields).encode() + value = int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") + return value % (corpus_length - (CONTEXT + 1)) + + +def train_window_start(seed: int, step: int, row: int, corpus_length: int) -> int: + payload = "\0".join( + [PROTOCOL_ID, "train-window", str(seed), str(step), str(row)] + ).encode() + value = int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") + return value % (corpus_length - (CONTEXT + 1)) + + +def concatenate_split(parquet_path: Path) -> tuple[bytes, int]: + table = pq.read_table(parquet_path, columns=["text"]) + rows = table.column("text").to_pylist() + payload = b"".join(((row or "") + "\n").encode("utf-8") for row in rows) + return payload, len(rows) + + +def main() -> None: + args = parse_args() + args.cache_dir.mkdir(parents=True, exist_ok=True) + + split_manifest: dict[str, Any] = {} + split_bytes: dict[str, bytes] = {} + for split in SPLITS: + relative = f"{DATASET_VARIANT}/{split}-00000-of-00001.parquet" + url = ( + f"https://huggingface.co/datasets/{DATASET_REPO}/resolve/" + f"{DATASET_REVISION}/{relative}" + ) + parquet_path = args.cache_dir / f"{split}.parquet" + download(url, parquet_path) + payload, rows = concatenate_split(parquet_path) + binary_path = args.cache_dir / f"{split}.bin" + if not binary_path.exists() or binary_path.read_bytes() != payload: + temporary = binary_path.with_suffix(".bin.tmp") + temporary.write_bytes(payload) + os.replace(temporary, binary_path) + split_bytes[split] = payload + split_manifest[split] = { + "source_path": relative, + "source_url": url, + "parquet_bytes": parquet_path.stat().st_size, + "parquet_sha256": file_sha256(parquet_path), + "rows": rows, + "concatenated_bytes": len(payload), + "concatenated_sha256": bytes_sha256(payload), + "binary_path": str(binary_path), + "binary_sha256": file_sha256(binary_path), + } + + train = split_bytes["train"] + validation = split_bytes["validation"] + + schedule_digest = hashlib.sha256() + schedule_cells = 0 + for seed in SEEDS: + for step in range(1, FORMAL_STEPS + 1): + for row in range(FORMAL_BATCH): + start = train_window_start(seed, step, row, len(train)) + schedule_digest.update(start.to_bytes(8, "big")) + schedule_cells += 1 + + validation_starts = [ + window_start("validation-window", index, len(validation)) + for index in range(VALIDATION_WINDOWS) + ] + diagnostic_starts = [ + window_start("diagnostic-window", index, len(validation)) + for index in range(DIAGNOSTIC_WINDOWS) + ] + + def tensor_hash(starts: list[int]) -> str: + digest = hashlib.sha256() + for start in starts: + digest.update(validation[start : start + CONTEXT + 1]) + return digest.hexdigest() + + manifest = { + "schema_version": 1, + "protocol_id": PROTOCOL_ID, + "status": "frozen-before-model-output", + "dataset": { + "repository": DATASET_REPO, + "revision": DATASET_REVISION, + "variant": DATASET_VARIANT, + "preprocessing": ( + "parquet row order; (text or empty string) + LF; UTF-8; " + "no normalization; vocabulary is raw bytes 0..255" + ), + "splits": split_manifest, + }, + "windows": { + "context": CONTEXT, + "target_bytes_per_window": CONTEXT, + "seeds": list(SEEDS), + "formal_steps": FORMAL_STEPS, + "formal_batch": FORMAL_BATCH, + "formal_schedule_cells": schedule_cells, + "formal_schedule_sha256": schedule_digest.hexdigest(), + "validation_starts": validation_starts, + "validation_tensor_sha256": tensor_hash(validation_starts), + "diagnostic_starts": diagnostic_starts, + "diagnostic_tensor_sha256": tensor_hash(diagnostic_starts), + }, + } + atomic_json(args.manifest, manifest) + print(json.dumps(manifest, ensure_ascii=False, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/experiments/k3/attnres/train.py b/experiments/k3/attnres/train.py new file mode 100644 index 0000000..7f55078 --- /dev/null +++ b/experiments/k3/attnres/train.py @@ -0,0 +1,715 @@ +#!/usr/bin/env python3 +"""Train one frozen residual variant for the reduced Attention Residuals study.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import math +import os +import platform +import statistics +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Iterable + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F + + +PROTOCOL_ID = "llm-atlas-k3-attnres-reduced-v1" +ARCHITECTURES = ("baseline", "full", "block") +EXPECTED_SEEDS = (2026073001, 2026073002, 2026073003) +EVAL_STEPS = (0, 100, 250, 500, 1000, 1500, 2000) +VOCABULARY = 256 +CONTEXT = 256 +LAYERS = 16 +SUBLAYERS = LAYERS * 2 +BLOCKS = 8 +SUBLAYERS_PER_BLOCK = SUBLAYERS // BLOCKS +D_MODEL = 192 +HEADS = 6 +D_HEAD = D_MODEL // HEADS +D_FF = 768 +RMS_EPS = 1e-6 +PEAK_LR = 3e-4 +MIN_LR = 3e-5 +WARMUP_STEPS = 100 +WEIGHT_DECAY = 0.1 +BETAS = (0.9, 0.95) +ADAM_EPS = 1e-8 +GRAD_CLIP = 1.0 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--architecture", choices=ARCHITECTURES, required=True) + parser.add_argument("--seed", type=int, required=True) + parser.add_argument("--steps", type=int, default=2000) + parser.add_argument("--batch-size", type=int, default=32) + parser.add_argument("--cache-dir", type=Path, required=True) + parser.add_argument("--manifest", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--validation-windows", type=int, default=64) + parser.add_argument("--diagnostic-windows", type=int, default=16) + parser.add_argument("--eval-batch-size", type=int, default=8) + parser.add_argument("--timing-warmup", type=int, default=20) + parser.add_argument("--run-kind", choices=("smoke", "formal", "replay"), default="formal") + return parser.parse_args() + + +def configure_determinism(seed: int) -> None: + if os.environ.get("CUBLAS_WORKSPACE_CONFIG") != ":4096:8": + raise RuntimeError("CUBLAS_WORKSPACE_CONFIG must be :4096:8 before Python starts") + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + torch.use_deterministic_algorithms(True) + torch.backends.cudnn.benchmark = False + torch.backends.cudnn.deterministic = True + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.set_float32_matmul_precision("highest") + + +def canonical_json_sha256(value: Any) -> str: + payload = json.dumps( + value, ensure_ascii=False, sort_keys=True, separators=(",", ":") + ).encode() + return hashlib.sha256(payload).hexdigest() + + +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 tensor_bytes(tensor: torch.Tensor) -> bytes: + value = tensor.detach().cpu().contiguous() + header = f"{value.dtype}|{tuple(value.shape)}|".encode() + return header + value.view(torch.uint8).numpy().tobytes() + + +def state_hash( + model: nn.Module, + *, + include_mixers: bool | None, +) -> str: + digest = hashlib.sha256() + for name, tensor in sorted(model.state_dict().items()): + is_mixer = name.startswith("mixers.") or name.startswith("output_mixer.") + if include_mixers is not None and is_mixer != include_mixers: + continue + digest.update(name.encode()) + digest.update(b"\0") + digest.update(tensor_bytes(tensor)) + return digest.hexdigest() + + +def window_start(seed: int, step: int, row: int, corpus_length: int) -> int: + payload = "\0".join( + [PROTOCOL_ID, "train-window", str(seed), str(step), str(row)] + ).encode() + value = int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") + return value % (corpus_length - (CONTEXT + 1)) + + +class ByteCorpus: + def __init__(self, cache_dir: Path, manifest: dict[str, Any], device: torch.device): + self.device = device + self.train = np.memmap(cache_dir / "train.bin", dtype=np.uint8, mode="r") + self.validation = np.memmap( + cache_dir / "validation.bin", dtype=np.uint8, mode="r" + ) + self.validation_starts = manifest["windows"]["validation_starts"] + self.diagnostic_starts = manifest["windows"]["diagnostic_starts"] + + def training_batch( + self, seed: int, step: int, batch_size: int + ) -> tuple[torch.Tensor, torch.Tensor]: + rows = np.empty((batch_size, CONTEXT + 1), dtype=np.int64) + for row in range(batch_size): + start = window_start(seed, step, row, len(self.train)) + rows[row] = self.train[start : start + CONTEXT + 1] + tensor = torch.from_numpy(rows).to(self.device, non_blocking=False) + return tensor[:, :-1], tensor[:, 1:] + + def fixed_batch( + self, starts: list[int], begin: int, end: int + ) -> tuple[torch.Tensor, torch.Tensor]: + chosen = starts[begin:end] + rows = np.empty((len(chosen), CONTEXT + 1), dtype=np.int64) + for row, start in enumerate(chosen): + rows[row] = self.validation[start : start + CONTEXT + 1] + tensor = torch.from_numpy(rows).to(self.device, non_blocking=False) + return tensor[:, :-1], tensor[:, 1:] + + +class RMSNorm(nn.Module): + def __init__(self, dimension: int): + super().__init__() + self.weight = nn.Parameter(torch.ones(dimension)) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + normalized = value.float() * torch.rsqrt( + value.float().square().mean(dim=-1, keepdim=True) + RMS_EPS + ) + return normalized.to(value.dtype) * self.weight + + +class CausalAttention(nn.Module): + def __init__(self): + super().__init__() + self.qkv = nn.Linear(D_MODEL, 3 * D_MODEL, bias=False) + self.o_proj = nn.Linear(D_MODEL, D_MODEL, bias=False) + mask = torch.triu(torch.ones(CONTEXT, CONTEXT, dtype=torch.bool), diagonal=1) + self.register_buffer("causal_mask", mask, persistent=False) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + batch, sequence, _ = value.shape + qkv = self.qkv(value).view(batch, sequence, 3, HEADS, D_HEAD) + query, key, content = qkv.unbind(dim=2) + query = query.transpose(1, 2) + key = key.transpose(1, 2) + content = content.transpose(1, 2) + scores = torch.matmul(query, key.transpose(-1, -2)).float() / math.sqrt(D_HEAD) + scores = scores.masked_fill( + self.causal_mask[:sequence, :sequence], float("-inf") + ) + probabilities = torch.softmax(scores, dim=-1).to(query.dtype) + mixed = torch.matmul(probabilities, content) + mixed = mixed.transpose(1, 2).contiguous().view(batch, sequence, D_MODEL) + return self.o_proj(mixed) + + +class SwiGLU(nn.Module): + def __init__(self): + super().__init__() + self.gate = nn.Linear(D_MODEL, D_FF, bias=False) + self.up = nn.Linear(D_MODEL, D_FF, bias=False) + self.down = nn.Linear(D_FF, D_MODEL, bias=False) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + return self.down(F.silu(self.gate(value)) * self.up(value)) + + +class TransformerBlock(nn.Module): + def __init__(self): + super().__init__() + self.attention_norm = RMSNorm(D_MODEL) + self.attention = CausalAttention() + self.mlp_norm = RMSNorm(D_MODEL) + self.mlp = SwiGLU() + + +class DepthMixer(nn.Module): + def __init__(self): + super().__init__() + self.query = nn.Parameter(torch.zeros(D_MODEL)) + self.key_norm = RMSNorm(D_MODEL) + + def forward( + self, sources: list[torch.Tensor], capture: bool = False + ) -> tuple[torch.Tensor, dict[str, Any] | None]: + values = torch.stack(sources, dim=0) + keys = self.key_norm(values) + logits = torch.einsum("d,nbtd->nbt", self.query, keys.float()) + weights = torch.softmax(logits, dim=0) + output = torch.einsum("nbt,nbtd->btd", weights, values.float()).to( + values.dtype + ) + if not capture: + return output, None + entropy = -(weights * torch.log(weights.clamp_min(1e-30))).sum(dim=0) + return output, { + "mean_weights": weights.mean(dim=(1, 2)).detach().cpu().tolist(), + "entropy_mean": entropy.mean().detach().cpu().item(), + "sources": len(sources), + } + + +@dataclass +class TraceAccumulator: + layer_input_rms: list[float] + branch_output_rms: list[float] + stream_state_rms: list[float] + depth_weights: list[dict[str, Any]] + output_weights: dict[str, Any] | None = None + + +def rms(value: torch.Tensor) -> float: + return value.float().square().mean().sqrt().detach().cpu().item() + + +class ReducedLanguageModel(nn.Module): + def __init__(self, architecture: str): + super().__init__() + self.architecture = architecture + self.token_embedding = nn.Embedding(VOCABULARY, D_MODEL) + self.position_embedding = nn.Embedding(CONTEXT, D_MODEL) + self.blocks = nn.ModuleList([TransformerBlock() for _ in range(LAYERS)]) + self.final_norm = RMSNorm(D_MODEL) + if architecture == "baseline": + self.mixers = nn.ModuleList() + self.output_mixer = None + else: + self.mixers = nn.ModuleList([DepthMixer() for _ in range(SUBLAYERS)]) + self.output_mixer = DepthMixer() + self.reset_parameters() + + def reset_parameters(self) -> None: + for module in self.modules(): + if isinstance(module, nn.Embedding): + nn.init.normal_(module.weight, mean=0.0, std=0.02) + elif isinstance(module, nn.Linear): + nn.init.normal_(module.weight, mean=0.0, std=0.02) + elif isinstance(module, RMSNorm): + nn.init.ones_(module.weight) + scaled = 0.02 / math.sqrt(2 * LAYERS) + for block in self.blocks: + nn.init.normal_(block.attention.o_proj.weight, mean=0.0, std=scaled) + nn.init.normal_(block.mlp.down.weight, mean=0.0, std=scaled) + for mixer in self.mixers: + nn.init.zeros_(mixer.query) + nn.init.ones_(mixer.key_norm.weight) + if self.output_mixer is not None: + nn.init.zeros_(self.output_mixer.query) + nn.init.ones_(self.output_mixer.key_norm.weight) + + def embed(self, input_ids: torch.Tensor) -> torch.Tensor: + positions = torch.arange(input_ids.shape[1], device=input_ids.device) + return self.token_embedding(input_ids) + self.position_embedding(positions)[None] + + def forward( + self, input_ids: torch.Tensor, capture: bool = False + ) -> tuple[torch.Tensor, TraceAccumulator | None]: + embedded = self.embed(input_ids) + trace = ( + TraceAccumulator([], [], [], []) + if capture + else None + ) + + if self.architecture == "baseline": + hidden = embedded + for block in self.blocks: + attention_input = hidden + attention_output = block.attention(block.attention_norm(attention_input)) + hidden = hidden + attention_output + if trace is not None: + trace.layer_input_rms.append(rms(attention_input)) + trace.branch_output_rms.append(rms(attention_output)) + trace.stream_state_rms.append(rms(hidden)) + mlp_input = hidden + mlp_output = block.mlp(block.mlp_norm(mlp_input)) + hidden = hidden + mlp_output + if trace is not None: + trace.layer_input_rms.append(rms(mlp_input)) + trace.branch_output_rms.append(rms(mlp_output)) + trace.stream_state_rms.append(rms(hidden)) + elif self.architecture == "full": + sources = [embedded] + mixer_index = 0 + for block in self.blocks: + attention_input, weights = self.mixers[mixer_index](sources, capture) + mixer_index += 1 + attention_output = block.attention(block.attention_norm(attention_input)) + sources.append(attention_output) + if trace is not None: + trace.layer_input_rms.append(rms(attention_input)) + trace.branch_output_rms.append(rms(attention_output)) + trace.stream_state_rms.append( + rms(torch.stack(sources, dim=0)) + ) + trace.depth_weights.append(weights or {}) + mlp_input, weights = self.mixers[mixer_index](sources, capture) + mixer_index += 1 + mlp_output = block.mlp(block.mlp_norm(mlp_input)) + sources.append(mlp_output) + if trace is not None: + trace.layer_input_rms.append(rms(mlp_input)) + trace.branch_output_rms.append(rms(mlp_output)) + trace.stream_state_rms.append( + rms(torch.stack(sources, dim=0)) + ) + trace.depth_weights.append(weights or {}) + assert self.output_mixer is not None + hidden, output_weights = self.output_mixer(sources, capture) + if trace is not None: + trace.output_weights = output_weights + else: + completed = [embedded] + partial: torch.Tensor | None = None + mixer_index = 0 + for block in self.blocks: + for branch_index in range(2): + sources = completed + ([] if partial is None else [partial]) + branch_input, weights = self.mixers[mixer_index](sources, capture) + mixer_index += 1 + if branch_index == 0: + branch_output = block.attention( + block.attention_norm(branch_input) + ) + else: + branch_output = block.mlp(block.mlp_norm(branch_input)) + partial = ( + branch_output if partial is None else partial + branch_output + ) + if trace is not None: + trace.layer_input_rms.append(rms(branch_input)) + trace.branch_output_rms.append(rms(branch_output)) + trace.stream_state_rms.append(rms(partial)) + trace.depth_weights.append(weights or {}) + if mixer_index % SUBLAYERS_PER_BLOCK == 0: + completed.append(partial) + partial = None + assert partial is None + assert len(completed) == BLOCKS + 1 + assert self.output_mixer is not None + hidden, output_weights = self.output_mixer(completed, capture) + if trace is not None: + trace.output_weights = output_weights + + normalized = self.final_norm(hidden) + logits = F.linear(normalized, self.token_embedding.weight) + return logits, trace + + +def learning_rate(step: int, total_steps: int) -> float: + if step <= WARMUP_STEPS: + return PEAK_LR * step / WARMUP_STEPS + progress = (step - WARMUP_STEPS) / max(1, total_steps - WARMUP_STEPS) + cosine = 0.5 * (1 + math.cos(math.pi * progress)) + return MIN_LR + (PEAK_LR - MIN_LR) * cosine + + +def cross_entropy(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: + return F.cross_entropy(logits.float().view(-1, VOCABULARY), targets.view(-1)) + + +@torch.no_grad() +def evaluate( + model: ReducedLanguageModel, + corpus: ByteCorpus, + starts: list[int], + window_count: int, + eval_batch_size: int, +) -> dict[str, float]: + model.eval() + loss_sum = 0.0 + target_count = 0 + for begin in range(0, window_count, eval_batch_size): + end = min(begin + eval_batch_size, window_count) + inputs, targets = corpus.fixed_batch(starts, begin, end) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + logits, _ = model(inputs) + loss = F.cross_entropy( + logits.float().view(-1, VOCABULARY), + targets.view(-1), + reduction="sum", + ) + loss_sum += loss.detach().cpu().item() + target_count += targets.numel() + nats = loss_sum / target_count + return {"cross_entropy_nats": nats, "bits_per_byte": nats / math.log(2)} + + +def percentile(values: list[float], quantile: float) -> float: + return float(np.quantile(np.asarray(values, dtype=np.float64), quantile)) + + +def core_parameter_gradient_rms(model: ReducedLanguageModel) -> list[float]: + values = [] + for block in model.blocks: + sum_square = 0.0 + count = 0 + for parameter in block.parameters(): + if parameter.grad is None: + continue + gradient = parameter.grad.detach().float() + sum_square += gradient.square().sum().detach().cpu().item() + count += gradient.numel() + values.append(math.sqrt(sum_square / count)) + return values + + +def diagnostic( + model: ReducedLanguageModel, + corpus: ByteCorpus, + window_count: int, +) -> dict[str, Any]: + model.eval() + model.zero_grad(set_to_none=True) + inputs, targets = corpus.fixed_batch( + corpus.diagnostic_starts, 0, window_count + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + logits, trace = model(inputs, capture=True) + loss = cross_entropy(logits, targets) + loss.backward() + gradients = core_parameter_gradient_rms(model) + assert trace is not None + return { + "loss_nats": loss.detach().cpu().item(), + "bits_per_byte": loss.detach().cpu().item() / math.log(2), + "layer_input_rms": trace.layer_input_rms, + "branch_output_rms": trace.branch_output_rms, + "stream_state_rms": trace.stream_state_rms, + "core_parameter_grad_rms_by_block": gradients, + "depth_weights": trace.depth_weights, + "output_weights": trace.output_weights, + } + + +def parameter_inventory(model: ReducedLanguageModel) -> dict[str, int]: + total = sum(parameter.numel() for parameter in model.parameters()) + mixer = sum( + parameter.numel() + for name, parameter in model.named_parameters() + if name.startswith("mixers.") or name.startswith("output_mixer.") + ) + embedding = model.token_embedding.weight.numel() + model.position_embedding.weight.numel() + return { + "total": total, + "core": total - mixer, + "mixer": mixer, + "embedding": embedding, + } + + +def main() -> None: + args = parse_args() + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required by the frozen protocol") + if args.run_kind != "smoke" and args.seed not in EXPECTED_SEEDS: + raise ValueError(f"formal/replay seed is not preregistered: {args.seed}") + configure_determinism(args.seed) + device = torch.device("cuda") + + manifest = json.loads(args.manifest.read_text()) + if manifest["protocol_id"] != PROTOCOL_ID: + raise ValueError("manifest protocol mismatch") + if manifest["dataset"]["revision"] != ( + "b08601e04326c79dfdd32d625aee71d232d685c3" + ): + raise ValueError("dataset revision mismatch") + corpus = ByteCorpus(args.cache_dir, manifest, device) + + model = ReducedLanguageModel(args.architecture).to(device) + initial_common_hash = state_hash(model, include_mixers=False) + initial_mixer_hash = ( + state_hash(model, include_mixers=True) + if args.architecture != "baseline" + else None + ) + inventory = parameter_inventory(model) + + decay_parameters: list[nn.Parameter] = [] + no_decay_parameters: list[nn.Parameter] = [] + for parameter in model.parameters(): + if parameter.ndim >= 2: + decay_parameters.append(parameter) + else: + no_decay_parameters.append(parameter) + optimizer = torch.optim.AdamW( + [ + {"params": decay_parameters, "weight_decay": WEIGHT_DECAY}, + {"params": no_decay_parameters, "weight_decay": 0.0}, + ], + lr=PEAK_LR, + betas=BETAS, + eps=ADAM_EPS, + ) + + evaluation_steps = sorted( + set(step for step in EVAL_STEPS if step <= args.steps) | {0, args.steps} + ) + evaluations = [ + { + "step": 0, + **evaluate( + model, + corpus, + corpus.validation_starts, + args.validation_windows, + args.eval_batch_size, + ), + } + ] + training_history: list[dict[str, float | int]] = [] + step_times: list[float] = [] + model.train() + + for step in range(1, args.steps + 1): + lr = learning_rate(step, args.steps) + for group in optimizer.param_groups: + group["lr"] = lr + inputs, targets = corpus.training_batch(args.seed, step, args.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, _ = model(inputs) + loss = cross_entropy(logits, targets) + if not torch.isfinite(loss): + raise RuntimeError(f"non-finite loss at step {step}: {loss}") + loss.backward() + unclipped_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), GRAD_CLIP + ) + optimizer.step() + torch.cuda.synchronize() + elapsed_ms = (time.perf_counter() - started) * 1000 + + if step == args.timing_warmup: + torch.cuda.reset_peak_memory_stats() + elif step > args.timing_warmup: + step_times.append(elapsed_ms) + + if step == 1 or step % 10 == 0 or step == args.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 evaluation_steps and step != 0: + evaluations.append( + { + "step": step, + **evaluate( + model, + corpus, + corpus.validation_starts, + args.validation_windows, + args.eval_batch_size, + ), + } + ) + model.train() + + training_peak_allocated = torch.cuda.max_memory_allocated() + training_peak_reserved = torch.cuda.max_memory_reserved() + diagnostic_result = diagnostic(model, corpus, args.diagnostic_windows) + final_common_hash = state_hash(model, include_mixers=False) + final_mixer_hash = ( + state_hash(model, include_mixers=True) + if args.architecture != "baseline" + else None + ) + timing = { + "warmup_steps_excluded": args.timing_warmup, + "measured_steps": len(step_times), + "mean_ms": statistics.fmean(step_times) if step_times else None, + "median_ms": statistics.median(step_times) if step_times else None, + "p95_ms": percentile(step_times, 0.95) if step_times else None, + "peak_allocated_bytes": training_peak_allocated, + "peak_reserved_bytes": training_peak_reserved, + } + + result = { + "schema_version": 1, + "protocol_id": PROTOCOL_ID, + "run_kind": args.run_kind, + "architecture": args.architecture, + "seed": args.seed, + "steps": args.steps, + "batch_size": args.batch_size, + "target_bytes_seen": args.steps * args.batch_size * CONTEXT, + "manifest": { + "path": str(args.manifest), + "file_sha256": file_sha256(args.manifest), + "formal_schedule_sha256": manifest["windows"][ + "formal_schedule_sha256" + ], + "validation_tensor_sha256": manifest["windows"][ + "validation_tensor_sha256" + ], + "diagnostic_tensor_sha256": manifest["windows"][ + "diagnostic_tensor_sha256" + ], + }, + "model": { + "layers": LAYERS, + "sublayers": SUBLAYERS, + "blocks_for_block_attnres": BLOCKS, + "sublayers_per_attnres_block": SUBLAYERS_PER_BLOCK, + "d_model": D_MODEL, + "heads": HEADS, + "d_head": D_HEAD, + "d_ff": D_FF, + "context": CONTEXT, + "vocabulary": VOCABULARY, + "parameters": inventory, + }, + "optimizer": { + "name": "AdamW", + "betas": list(BETAS), + "epsilon": ADAM_EPS, + "weight_decay_ndim_ge_2": WEIGHT_DECAY, + "peak_lr": PEAK_LR, + "min_lr": MIN_LR, + "warmup_steps": WARMUP_STEPS, + "grad_clip": GRAD_CLIP, + }, + "hashes": { + "initial_common_parameters": initial_common_hash, + "initial_mixer_parameters": initial_mixer_hash, + "final_common_parameters": final_common_hash, + "final_mixer_parameters": final_mixer_hash, + }, + "evaluations": evaluations, + "training_history": training_history, + "diagnostic": diagnostic_result, + "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", + "compile": False, + }, + } + result["canonical_sha256_without_self"] = canonical_json_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), + "architecture": args.architecture, + "seed": args.seed, + "steps": args.steps, + "final_bpc": evaluations[-1]["bits_per_byte"], + "initial_common_hash": initial_common_hash, + "final_common_hash": final_common_hash, + "canonical_sha256": result["canonical_sha256_without_self"], + "timing": timing, + }, + ensure_ascii=False, + indent=2, + ) + ) + + +if __name__ == "__main__": + main()