research: add reduced AttnRes runner

This commit is contained in:
wuyang
2026-07-30 06:38:26 +08:00
parent e3617532de
commit f99872039f
3 changed files with 950 additions and 0 deletions
+43
View File
@@ -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.
+192
View File
@@ -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()
+715
View File
@@ -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()