code-fdsp-v2 / App_ddp.py
Bc-AI's picture
Upload App_ddp.py
7ed8156 verified
Raw History Blame Contribute Delete
177 kB
"""
Orion Flagship 2B — T2.2 DDP base-pretraining trainer.
Target hardware:
One node with eight NVIDIA H200s (SM 90, ~141 GB/GPU).
Blackwell SM 120 is also accepted, subject to memory/startup checks.
T2.2 changes over the original Mini prototype:
- ~2.042B total parameters at vocab size 32,768 (2,041,951,637;
including sparsely activated experts, independent of GPU hardware).
- Smaller Knowledge Vault capacity reallocates parameters to wider
attention, specialist experts, and the shared expert.
- Procedure Bank v2: top-1 specialist, small shared expert, soft NULL gate,
and routing conditioned on Working State + stage + deliberation cycle.
- Knowledge Vault v2: state/stage-aware product-key query, confidence-aware
grouped gating, and learned residual-delta fusion.
- Working State retains the bounded convex update that was stable in T2.
- Router-specialization objective warms in slowly rather than being active
at full strength from step zero.
- Production telemetry, gradient-health checks, warnings, and periodic
causal ablation probes catch modules that silently become decorative.
Data:
Pure causal base pretraining only. Expects the completed 50B-token
orion-nano-base-v1 cache produced by orion_nano_base_pretokenizer_v1.py.
Safety / resumability:
- Automatically starts fresh only when remote checkpoint discovery succeeds
and confirms that no complete checkpoint exists.
- Otherwise newest complete local/remote checkpoint wins.
- Dataset revision and deterministic stratified data cursor are frozen inside checkpoints.
- HF_TOKEN is read from the environment; no credentials are embedded.
- Automatic destructive Hub history squashing is disabled by default.
This is a new architecture/run. Nano/FSDP and prior T2.1 DDP checkpoints are incompatible.
Distributed launch (one node with eight supported GPUs):
torchrun --standalone --nproc_per_node=8 App_ddp.py
The global micro-batch is 16 sequences (2/GPU) with 7 accumulation steps,
preserving 114,688 tokens/update. DDP replicates the full model and Adam states
on each GPU; memory is NOT pooled across GPUs. Checkpoints use format 4,
ddp_full_state. Older checkpoint formats are intentionally rejected.
"""
import os
import re
import sys
import importlib.util
import subprocess
# ===========================================================================
# SETTINGS — edit these here, not in environment variables
# ===========================================================================
MODEL_NAME = "Orion Flagship 2B T2.2"
# Separate checkpoint destination for the incompatible 2B run.
HF_REPO_ID = os.environ.get(
"ORION_2B_MODEL_REPO",
"Project-Prism/Orion-Flagship-2B-T2.2",
)
HF_PRIVATE = True
HF_TOKEN = os.environ.get("HF_TOKEN", "").strip()
# 50B-token BASE corpus built by orion_nano_base_pretokenizer_v1.py.
DATA_REPO_ID = "smilyai-large-team/nano-orion-ultra"
DATA_REPO_TYPE = "dataset"
DATA_REVISION = "main" # Resolved to one immutable commit for a new stream.
DATA_CACHE_DIR = "nano_base_token_cache_v1"
DATA_MANIFEST = f"{DATA_CACHE_DIR}/manifest.json"
DATA_TARGET_TOKENS = 50_000_000_000
TOKENIZER_NAME = "mistralai/Mistral-7B-v0.3"
WORK_DIR = os.environ.get("ORION_2B_WORK_DIR", "./orion_2b_t22_ddp_work")
# Fresh/resume policy. False is safe: choose_resume auto-starts only if Hub
# discovery succeeds and finds no complete checkpoint at all.
RESUME_LOCAL_PATH = None
START_FROM_SCRATCH = False
ALLOW_FRESH_START_WITH_EXISTING_REMOTE = False
ALLOW_REMOTE_POINTER_ROLLBACK = False
# Conservative 2B replicated-model batch for H200; keep 1K context.
SEQUENCE_LENGTH = 1024
GLOBAL_MICRO_BATCH_SIZE = 16
# For OOMs, set this to 8: 1 sequence/GPU × 8 GPUs × 14 accumulation steps.
# Set to the per-rank value after torch.distributed is initialized. Keeping
# the global value explicit prevents accidentally multiplying the update batch
# by eight when launching with torchrun.
MICRO_BATCH_SIZE = GLOBAL_MICRO_BATCH_SIZE
TOKENS_PER_UPDATE = 114_688
# Nominal corpus-pass budget, not exact per-example epochs: ShardStream replays
# each source with deterministic reshuffling, while omitting existing shard tails.
# Leave the final partial optimizer update unused instead of exceeding the budget.
TRAIN_EPOCHS = 2
REQUESTED_TRAIN_TOKENS = DATA_TARGET_TOKENS * TRAIN_EPOCHS
TRAIN_UPDATE_COUNT = REQUESTED_TRAIN_TOKENS // TOKENS_PER_UPDATE
TRAIN_TOKENS = TRAIN_UPDATE_COUNT * TOKENS_PER_UPDATE
MAX_STEPS = TRAIN_UPDATE_COUNT
WARMUP_STEPS = 2_000
PEAK_LR = 2.0e-4
MIN_LR = 2.0e-5
WEIGHT_DECAY = 0.1
GRAD_CLIP = 1.0
ROUTER_AUX_COEF = 0.005
ROUTER_Z_COEF = 0.0005
ROUTER_SPECIALIZATION_COEF = 0.0005
SPEC_WARMUP_STEPS = 5_000
# Straight-through-style task gradient into the chosen top-1 probability.
# Forward scale stays ~1 while the router still receives a bounded task signal.
ROUTER_TASK_GRAD_SCALE = 0.10
UNCHECKPOINTED_LAST_N = 2
CHECKPOINT_DELIBERATION = True
CHECKPOINT_LOSS_CHUNKS = False
LOSS_CHUNK_TOKENS = 1024
COMPILE_ATTENTION = True
USE_FUSED_MOE_COMBINE = True
USE_FUSED_MOE_PACK = True
USE_FUSED_CROSS_ENTROPY = True
USE_GROUPED_EXPERT_GEMM = False # v2 top-1 path uses the proven pack/combine route.
BENCHMARK_EXPERT_PATHS = False
RUN_KERNEL_TESTS = True
GROUPED_EXPERT_GEMM_ACTIVE = False
AUTOCAST_CACHE = True
# The fused kernels and attention path are intentionally gated to the two
# architectures targeted by this run. Hardware execution must still be checked.
# Do not silently run on a nearby
# compute capability: Triton/PTX support and BF16/Flash behavior can differ.
SUPPORTED_COMPUTE_CAPABILITIES = ((9, 0), (12, 0))
HARDWARE_PROFILE_NAMES = {
(9, 0): "Hopper / H200",
(12, 0): "Blackwell",
}
KERNEL_TEST_TIMEOUT_SECONDS = 120.0
ATTENTION_PROBE_TIMEOUT_SECONDS = 180.0
CPU_THREADS = 8
PREFETCH_BATCHES = 4
DATA_SEED = 42
# Data-stream v3: source-stratified interleaving.
# Every 25 training sequences contains exactly the frozen 60/12/8/20 source mix,
# in a deterministic shuffled order. This prevents 268M-token single-domain runs.
DATA_STREAM_VERSION = 3
RUN_ID = "orion-2b-t22-ddp-2epochs-v5"
MIX_CYCLE = (
["general_fineweb_edu"] * 15
+ ["math_finemath_4plus"] * 3
+ ["math_openwebmath"] * 2
+ ["code_python_clean"] * 5
)
MAPPED_SHARD_CACHE = 8
DOWNLOAD_WORKERS = 4
LOG_EVERY_STEPS = 10
T2_TELEMETRY_EVERY_STEPS = 100
T2_ABLATION_EVERY_STEPS = 1_000
ABLATION_PROBE_BATCH = 2
STARTUP_MODEL_PROBE = True
CHECKPOINT_EVERY_MINUTES = 60
SESSION_HOURS = 12.0
MAX_TRAIN_HOURS = 11.0
UPLOAD_RESERVE_MINUTES = 30
SAVE_RESERVE_MINUTES = 5
FINAL_SAVE_GUARD_MINUTES = 30
ESTIMATED_UPLOAD_MB_PER_SECOND = 30.0
USE_LARGE_FOLDER_UPLOAD = True
UPLOAD_WORKERS = 8
KEEP_ONLY_LATEST_REMOTE_FOLDER = True
# Off by default: history rewriting is destructive. Enable manually only on a
# dedicated rolling-checkpoint repo if storage pressure makes it necessary.
SUPER_SQUASH_AFTER_REMOTE_CLEANUP = False
HF_API_MAX_RETRIES = 8
HF_API_RETRY_BASE_SECONDS = 5.0
HF_API_RETRY_MAX_SECONDS = 120.0
MAX_CONSECUTIVE_NONFINITE = 10
AUTO_INSTALL_MISSING = True
# Explicitly bounded validation path. Production training is unchanged unless
# this is enabled in the launch environment, for example:
# ORION_2B_PREFLIGHT_SMOKE=1 torchrun --standalone --nproc_per_node=8 App_ddp.py
PREFLIGHT_SMOKE = os.environ.get(
"ORION_2B_PREFLIGHT_SMOKE", os.environ.get("ORION_NANO_PREFLIGHT_SMOKE", "")
).strip().lower() in {
"1", "true", "yes", "on",
}
if PREFLIGHT_SMOKE:
# Keep the synthetic checkpoint and its metadata out of the production
# resume directory when both modes share ORION_2B_WORK_DIR.
WORK_DIR = os.path.join(WORK_DIR, "preflight-smoke")
# ===========================================================================
# Internal runtime configuration — no external configuration needed
# ===========================================================================
if MICRO_BATCH_SIZE < 1:
raise ValueError("MICRO_BATCH_SIZE must be positive.")
if UNCHECKPOINTED_LAST_N < 0:
raise ValueError("UNCHECKPOINTED_LAST_N must be non-negative.")
micro_tokens = MICRO_BATCH_SIZE * SEQUENCE_LENGTH
if TOKENS_PER_UPDATE % micro_tokens:
raise ValueError(
"MICRO_BATCH_SIZE * SEQUENCE_LENGTH must divide TOKENS_PER_UPDATE."
)
GRAD_ACCUM_STEPS = TOKENS_PER_UPDATE // micro_tokens
if LOSS_CHUNK_TOKENS is not None and LOSS_CHUNK_TOKENS < 1:
raise ValueError("LOSS_CHUNK_TOKENS must be positive or None.")
if ESTIMATED_UPLOAD_MB_PER_SECOND <= 0:
raise ValueError("ESTIMATED_UPLOAD_MB_PER_SECOND must be positive.")
# These are internal library settings, not user-facing environment options.
# They must be established before importing torch/triton.
os.environ["TRITON_CACHE_DIR"] = os.path.join(WORK_DIR, "triton_cache")
os.environ["TORCHINDUCTOR_CACHE_DIR"] = os.path.join(
WORK_DIR, "inductor_cache"
)
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
os.environ["TOKENIZERS_PARALLELISM"] = "false"
# torchrun gives every worker its own compiler cache. Sharing these caches
# between eight compiler processes can corrupt artifacts on some filesystems.
_LOCAL_RANK_ENV = os.environ.get("LOCAL_RANK", "0")
os.environ["TRITON_CACHE_DIR"] = os.path.join(
WORK_DIR, "triton_cache", f"rank-{_LOCAL_RANK_ENV}"
)
os.environ["TORCHINDUCTOR_CACHE_DIR"] = os.path.join(
WORK_DIR, "inductor_cache", f"rank-{_LOCAL_RANK_ENV}"
)
def install_missing():
requirements = {
"torch": "torch",
"triton": "triton",
"numpy": "numpy",
"transformers": "transformers>=4.48",
"huggingface_hub": "huggingface_hub>=0.32",
"safetensors": "safetensors",
"sentencepiece": "sentencepiece",
}
missing = [
requirement
for module, requirement in requirements.items()
if importlib.util.find_spec(module) is None
]
if missing:
if not AUTO_INSTALL_MISSING:
raise RuntimeError(f"Install missing packages: {missing}")
subprocess.check_call([
sys.executable,
"-m",
"pip",
"install",
"--no-cache-dir",
*missing,
])
install_missing()
import copy
import datetime
import gc
import hashlib
import json
import math
import queue
import random
import shutil
import signal
import threading
import time
import uuid
from collections import OrderedDict
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path, PurePosixPath
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
import triton
import triton.language as tl
from torch.nn.parallel import DistributedDataParallel as DDP
from huggingface_hub import HfApi, hf_hub_download, snapshot_download
from huggingface_hub.errors import EntryNotFoundError
from safetensors.torch import save_file, load_file
from torch.nn.attention import sdpa_kernel, SDPBackend
from torch.utils.checkpoint import checkpoint
from transformers import AutoTokenizer
torch.set_num_threads(min(CPU_THREADS, os.cpu_count() or 1))
WORK = Path(WORK_DIR)
STOP_REQUESTED = False
RANK = int(os.environ.get("RANK", "0"))
LOCAL_RANK = int(os.environ.get("LOCAL_RANK", "0"))
WORLD_SIZE = int(os.environ.get("WORLD_SIZE", "1"))
DEVICE = torch.device("cuda", LOCAL_RANK)
def rank0_print(*args, **kwargs):
if RANK == 0:
print(*args, **kwargs)
def _format_capability(capability):
return f"SM {capability[0]}.{capability[1]}"
def _hardware_inventory():
"""Return visible GPU details without making a profile assumption."""
inventory = []
for index in range(torch.cuda.device_count()):
properties = torch.cuda.get_device_properties(index)
capability = tuple(torch.cuda.get_device_capability(index))
inventory.append({
"index": index,
"name": properties.name,
"memory_gib": properties.total_memory / 1024**3,
"capability": capability,
})
return inventory
def _format_hardware_inventory(inventory):
return "\n".join(
f" GPU {item['index']}: {item['name']} | "
f"{_format_capability(item['capability'])} | "
f"{item['memory_gib']:.1f} GiB"
for item in inventory
) or " (no visible CUDA GPUs)"
def validate_hardware():
"""Validate the visible rank devices before any real data is touched."""
if not torch.cuda.is_available():
raise RuntimeError("CUDA GPU required; torch.cuda.is_available() is false.")
inventory = _hardware_inventory()
visible_count = len(inventory)
if visible_count < WORLD_SIZE:
raise RuntimeError(
f"GPU count mismatch: torchrun requested {WORLD_SIZE} ranks but "
f"only {visible_count} CUDA GPU(s) are visible.\n"
f"Visible devices:\n{_format_hardware_inventory(inventory)}"
)
if LOCAL_RANK >= visible_count:
raise RuntimeError(
f"LOCAL_RANK={LOCAL_RANK} is outside the {visible_count} visible "
"CUDA GPU(s).\n"
f"Visible devices:\n{_format_hardware_inventory(inventory)}"
)
participating = inventory[:WORLD_SIZE]
unsupported = [
item for item in participating
if item["capability"] not in SUPPORTED_COMPUTE_CAPABILITIES
]
if unsupported:
supported = ", ".join(
_format_capability(capability)
for capability in SUPPORTED_COMPUTE_CAPABILITIES
)
raise RuntimeError(
"Unsupported CUDA compute capability for this run. Supported "
f"capabilities are {supported}; every participating rank must use "
"one of them.\n"
f"Visible devices:\n{_format_hardware_inventory(inventory)}"
)
capabilities = {item["capability"] for item in participating}
if len(capabilities) != 1:
raise RuntimeError(
"Mixed GPU compute capabilities are not supported for one DDP run; "
"use homogeneous H200/SM 90 or Blackwell/SM 120 ranks.\n"
f"Visible devices:\n{_format_hardware_inventory(inventory)}"
)
return inventory[LOCAL_RANK], inventory
def dist_barrier():
if dist.is_initialized():
dist.barrier()
def broadcast_object(value):
if not dist.is_initialized():
return value
values = [value if RANK == 0 else None]
dist.broadcast_object_list(values, src=0)
return values[0]
# ===========================================================================
# Utilities
# ===========================================================================
def safe_relative_path(value):
text = str(value)
path = PurePosixPath(text)
if (
not text
or path.is_absolute()
or ".." in path.parts
or "\\" in text
or not path.parts
):
raise ValueError(f"Unsafe relative path: {value!r}")
return str(path)
def atomic_json(path, value):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_name(path.name + "." + uuid.uuid4().hex + ".tmp")
with tmp.open("w") as f:
json.dump(value, f, indent=2)
f.flush()
os.fsync(f.fileno())
os.replace(tmp, path)
def nested_bytes(value):
if torch.is_tensor(value):
return value.numel() * value.element_size()
if isinstance(value, dict):
return sum(nested_bytes(v) for v in value.values())
if isinstance(value, (tuple, list)):
return sum(nested_bytes(v) for v in value)
return 0
def cpu_tree(value):
# Pageable checkpoint copies, not large pinned allocations.
if torch.is_tensor(value):
return value.detach().to("cpu", copy=True).contiguous()
if isinstance(value, dict):
return {k: cpu_tree(v) for k, v in value.items()}
if isinstance(value, list):
return [cpu_tree(v) for v in value]
if isinstance(value, tuple):
return tuple(cpu_tree(v) for v in value)
return value
def cuda_tree(value):
if torch.is_tensor(value):
return value.to(DEVICE)
if isinstance(value, dict):
return {k: cuda_tree(v) for k, v in value.items()}
if isinstance(value, list):
return [cuda_tree(v) for v in value]
if isinstance(value, tuple):
return tuple(cuda_tree(v) for v in value)
return value
def shard_items(items, target_bytes=512 * 1024**2):
shard, size = {}, 0
for name, value in items:
n = nested_bytes(value)
if shard and size + n > target_bytes:
yield shard
shard, size = {}, 0
shard[name] = value
size += n
if shard:
yield shard
def capture_rng():
return {
"python_rng": random.getstate(),
"numpy_rng": np.random.get_state(),
"torch_rng": torch.get_rng_state(),
"cuda_rng": torch.cuda.get_rng_state(DEVICE),
}
def restore_rng(state):
random.setstate(state["python_rng"])
if "numpy_rng" in state:
np.random.set_state(state["numpy_rng"])
torch.set_rng_state(state["torch_rng"])
torch.cuda.set_rng_state(state["cuda_rng"], device=DEVICE)
def folder_bytes(path):
return sum(
p.stat().st_size for p in Path(path).rglob("*") if p.is_file()
)
# ===========================================================================
# Triton RMSNorm
# ===========================================================================
@triton.jit
def rms_forward_kernel(
X, Y, INV,
D: tl.constexpr,
EPS: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0)
col = tl.arange(0, BLOCK)
x = tl.load(
X + row * D + col, mask=col < D, other=0
).to(tl.float32)
inv = tl.rsqrt(tl.sum(x * x, axis=0) / D + EPS)
tl.store(Y + row * D + col, x * inv, mask=col < D)
tl.store(INV + row, inv)
@triton.jit
def rms_backward_kernel(
X, DY, INV, DX,
D: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0)
col = tl.arange(0, BLOCK)
mask = col < D
x = tl.load(X + row * D + col, mask, other=0).to(tl.float32)
dy = tl.load(DY + row * D + col, mask, other=0).to(tl.float32)
inv = tl.load(INV + row)
projection = tl.sum(x * dy, axis=0) / D
dx = inv * dy - x * inv * inv * inv * projection
tl.store(DX + row * D + col, dx, mask)
class TritonRMSFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x, eps):
x = x.contiguous()
d = x.shape[-1]
rows = x.numel() // d
y = torch.empty_like(x)
inv = torch.empty(rows, device=x.device, dtype=torch.float32)
rms_forward_kernel[(rows,)](
x, y, inv,
D=d,
EPS=eps,
BLOCK=triton.next_power_of_2(d),
)
ctx.save_for_backward(x, inv)
return y
@staticmethod
def backward(ctx, dy):
x, inv = ctx.saved_tensors
dy = dy.contiguous()
d = x.shape[-1]
dx = torch.empty_like(x)
rms_backward_kernel[(x.numel() // d,)](
x, dy, inv, dx,
D=d,
BLOCK=triton.next_power_of_2(d),
)
return dx, None
class RMSNorm(nn.Module):
def forward(self, x):
return TritonRMSFunction.apply(x, 1e-6)
# ===========================================================================
# Triton SwiGLU
# ===========================================================================
@triton.jit
def swiglu_forward_kernel(A, B, Y, N, BLOCK: tl.constexpr):
idx = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = idx < N
a = tl.load(A + idx, mask, other=0).to(tl.float32)
b = tl.load(B + idx, mask, other=0).to(tl.float32)
s = 1.0 / (1.0 + tl.exp(-a))
tl.store(Y + idx, a * s * b, mask)
@triton.jit
def swiglu_backward_kernel(
A, B, DY, DA, DB, N, BLOCK: tl.constexpr
):
idx = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = idx < N
a = tl.load(A + idx, mask, other=0).to(tl.float32)
b = tl.load(B + idx, mask, other=0).to(tl.float32)
dy = tl.load(DY + idx, mask, other=0).to(tl.float32)
s = 1.0 / (1.0 + tl.exp(-a))
tl.store(DA + idx, dy * b * s * (1.0 + a * (1.0 - s)), mask)
tl.store(DB + idx, dy * a * s, mask)
class TritonSwiGLUFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, a, b):
a, b = a.contiguous(), b.contiguous()
y = torch.empty_like(a)
swiglu_forward_kernel[(triton.cdiv(a.numel(), 1024),)](
a, b, y, a.numel(), BLOCK=1024, num_warps=4
)
ctx.save_for_backward(a, b)
return y
@staticmethod
def backward(ctx, dy):
a, b = ctx.saved_tensors
dy = dy.contiguous()
da, db = torch.empty_like(a), torch.empty_like(b)
swiglu_backward_kernel[(triton.cdiv(a.numel(), 1024),)](
a, b, dy, da, db, a.numel(), BLOCK=1024, num_warps=4
)
return da, db
# ===========================================================================
# Fused MoE combine and backward
# ===========================================================================
@triton.jit
def moe_combine_forward_kernel(
GROUPED, WEIGHTS, INVERSE, OUTPUT,
D: tl.constexpr,
K: tl.constexpr,
BLOCK: tl.constexpr,
):
token = tl.program_id(0)
col = tl.arange(0, BLOCK)
mask = col < D
result = tl.full((BLOCK,), 0.0, tl.float32)
for slot in tl.static_range(K):
assignment = token * K + slot
grouped_row = tl.load(INVERSE + assignment)
weight = tl.load(WEIGHTS + assignment).to(tl.float32)
value = tl.load(
GROUPED + grouped_row * D + col, mask, other=0
).to(tl.float32)
result = result + value * weight
tl.store(OUTPUT + token * D + col, result, mask)
@triton.jit
def moe_combine_backward_kernel(
GROUPED, WEIGHTS, INVERSE,
DOUTPUT, DGROUPED, DWEIGHTS,
D: tl.constexpr,
K: tl.constexpr,
BLOCK: tl.constexpr,
):
token = tl.program_id(0)
col = tl.arange(0, BLOCK)
mask = col < D
dy = tl.load(
DOUTPUT + token * D + col, mask, other=0
).to(tl.float32)
for slot in tl.static_range(K):
assignment = token * K + slot
grouped_row = tl.load(INVERSE + assignment)
weight = tl.load(WEIGHTS + assignment).to(tl.float32)
value = tl.load(
GROUPED + grouped_row * D + col, mask, other=0
).to(tl.float32)
# INVERSE is a permutation: each destination row is unique.
tl.store(
DGROUPED + grouped_row * D + col,
dy * weight,
mask,
)
tl.store(
DWEIGHTS + assignment,
tl.sum(dy * value, axis=0),
)
class FusedMoECombine(torch.autograd.Function):
@staticmethod
def forward(ctx, grouped, weights, inverse):
grouped = grouped.contiguous()
weights = weights.contiguous()
inverse = inverse.contiguous()
n, k = weights.shape
d = grouped.shape[1]
if grouped.shape[0] != n * k or inverse.numel() != n * k:
raise ValueError("Invalid MoE combine shapes.")
if weights.dtype != torch.float32:
raise ValueError("MoE routing weights must be FP32.")
output = torch.empty(
(n, d), device=grouped.device, dtype=torch.float32
)
moe_combine_forward_kernel[(n,)](
grouped, weights, inverse, output,
D=d,
K=k,
BLOCK=triton.next_power_of_2(d),
num_warps=8 if d >= 2048 else 4,
enable_fp_fusion=False,
)
ctx.save_for_backward(grouped, weights, inverse)
return output
@staticmethod
def backward(ctx, grad_output):
grouped, weights, inverse = ctx.saved_tensors
grad_output = grad_output.contiguous()
n, k = weights.shape
d = grouped.shape[1]
grad_grouped = torch.empty_like(grouped)
grad_weights = torch.empty_like(weights)
moe_combine_backward_kernel[(n,)](
grouped, weights, inverse,
grad_output, grad_grouped, grad_weights,
D=d,
K=k,
BLOCK=triton.next_power_of_2(d),
num_warps=8 if d >= 2048 else 4,
enable_fp_fusion=False,
)
return grad_grouped, grad_weights, None
def combine_reference(grouped, weights, inverse):
n, k = weights.shape
d = grouped.shape[-1]
selected = grouped.index_select(0, inverse).view(n, k, d).float()
return (selected * weights.unsqueeze(-1)).sum(dim=1)
@triton.jit
def moe_pack_forward_kernel(X, ORDER, PACKED, D: tl.constexpr,
K: tl.constexpr, BLOCK: tl.constexpr):
row = tl.program_id(0)
col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
token = tl.load(ORDER + row) // K
value = tl.load(X + token * D + col, col < D, other=0)
# Gather and autocast in one pass, without a full-size cast temporary.
tl.store(PACKED + row * D + col, value, col < D)
@triton.jit
def moe_pack_backward_kernel(DPACKED, INVERSE, DX, D: tl.constexpr,
K: tl.constexpr, BLOCK: tl.constexpr):
token = tl.program_id(0)
col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
value = tl.full((BLOCK,), 0.0, tl.float32)
for slot in tl.static_range(K):
row = tl.load(INVERSE + token * K + slot)
grad = tl.load(DPACKED + row * D + col, col < D, other=0)
value += grad.to(tl.float32)
# Match the index_select gradient's compute-dtype buffer before the
# gradient passes back through the original FP32 -> BF16 cast.
value = value.to(DPACKED.dtype.element_ty).to(tl.float32)
tl.store(DX + token * D + col, value, col < D)
class FusedMoEPack(torch.autograd.Function):
@staticmethod
def forward(ctx, x, order, inverse, k, dtype):
x = x.contiguous()
n, d = x.shape
packed = torch.empty((n * k, d), device=x.device, dtype=dtype)
moe_pack_forward_kernel[(n * k, triton.cdiv(d, 512))](
x, order, packed, D=d, K=k, BLOCK=512, num_warps=4,
)
ctx.save_for_backward(inverse)
ctx.input_shape, ctx.input_dtype, ctx.k = x.shape, x.dtype, k
return packed
@staticmethod
def backward(ctx, grad):
inverse, = ctx.saved_tensors
grad = grad.contiguous()
n, d = ctx.input_shape
dx = torch.empty(ctx.input_shape, device=grad.device,
dtype=ctx.input_dtype)
moe_pack_backward_kernel[(n, triton.cdiv(d, 512))](
grad, inverse, dx, D=d, K=ctx.k, BLOCK=512, num_warps=4,
)
return dx, None, None, None, None
@triton.jit
def cross_entropy_forward_kernel(LOGITS, TARGETS, LOSS, LSE,
V: tl.constexpr, BLOCK: tl.constexpr):
row = tl.program_id(0)
col = tl.arange(0, BLOCK)
x = tl.load(LOGITS + row * V + col, col < V,
other=-float("inf")).to(tl.float32)
maximum = tl.max(x, 0)
log_sum = tl.log(tl.sum(tl.exp(x - maximum), 0))
target = tl.load(TARGETS + row)
chosen = tl.load(LOGITS + row * V + target,
(target >= 0) & (target < V), other=0).to(tl.float32)
# Subtract before adding log_sum to avoid large-logit cancellation.
loss = (maximum - chosen) + log_sum
tl.store(LOSS + row, tl.where(target == -100, 0.0, loss))
tl.store(LSE + row * 2, maximum)
tl.store(LSE + row * 2 + 1, log_sum)
@triton.jit
def cross_entropy_backward_kernel(LOGITS, TARGETS, LSE, DLOSS, DLOGITS,
V: tl.constexpr, BLOCK: tl.constexpr):
row = tl.program_id(0)
col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
x = tl.load(LOGITS + row * V + col, col < V, other=0).to(tl.float32)
maximum = tl.load(LSE + row * 2)
log_sum = tl.load(LSE + row * 2 + 1)
target = tl.load(TARGETS + row)
upstream = tl.load(DLOSS)
prob = tl.exp((x - maximum) - log_sum)
dx = (prob - (col == target).to(tl.float32)) * upstream
tl.store(DLOGITS + row * V + col,
tl.where(target == -100, 0.0, dx), col < V)
class FusedCrossEntropy(torch.autograd.Function):
"""Unweighted sum CE; FP32 reductions, original logits dtype gradient.
Reads BF16 logits directly, avoiding autocast's full FP32 logits and
log-softmax buffers. Does not overwrite logits saved by checkpointing.
"""
@staticmethod
def forward(ctx, logits, targets):
logits, targets = logits.contiguous(), targets.contiguous()
n, v = logits.shape
losses = torch.empty(n, device=logits.device, dtype=torch.float32)
lse = torch.empty((n, 2), device=logits.device, dtype=torch.float32)
cross_entropy_forward_kernel[(n,)](
logits, targets, losses, lse, V=v,
BLOCK=triton.next_power_of_2(v), num_warps=16 if v >= 16384 else 4,
enable_fp_fusion=False,
)
ctx.save_for_backward(logits, targets, lse)
return losses.sum()
@staticmethod
def backward(ctx, grad):
logits, targets, lse = ctx.saved_tensors
n, v = logits.shape
dx = torch.empty_like(logits)
cross_entropy_backward_kernel[(n, triton.cdiv(v, 1024))](
logits, targets, lse, grad, dx, V=v, BLOCK=1024,
num_warps=4, enable_fp_fusion=False,
)
return dx, None
# ===========================================================================
# Kernel checks
# ===========================================================================
def _check_preflight_deadline(deadline, label):
if deadline is not None and time.monotonic() >= deadline:
raise TimeoutError(
f"{label} exceeded its preflight time limit; "
"reduce the smoke-test scope or inspect the CUDA/Triton setup."
)
def test_optimized_kernels(device="cuda", dtypes=None, deadline=None):
# Test odd tails, routing collisions, cast-backward rounding, ignored
# labels, large logits, and the actual vocabulary-sized reduction.
if dtypes is None:
dtypes = (torch.float32, torch.bfloat16)
for dtype in dtypes:
for d in (127, 2048):
_check_preflight_deadline(deadline, "Triton kernel preflight")
n, k = 17, 2
x = torch.randn(n, d, device=device, requires_grad=True)
order = torch.randperm(n * k, device=device)
inverse = torch.argsort(order)
upstream = torch.randn(n * k, d, device=device, dtype=dtype)
actual = FusedMoEPack.apply(x, order, inverse, k, dtype)
reference = x.to(dtype).index_select(0, order // k)
dx = torch.autograd.grad(actual, x, upstream)[0]
dxr = torch.autograd.grad(reference, x, upstream)[0]
torch.testing.assert_close(actual, reference, rtol=0, atol=0)
torch.testing.assert_close(dx, dxr, rtol=0, atol=0)
for v in (127, 32768):
_check_preflight_deadline(deadline, "Triton kernel preflight")
logits = (torch.randn(9, v, device=device) * 4 + 80).to(dtype)
logits.requires_grad_(True)
targets = torch.randint(v, (9,), device=device)
targets[0] = -100
actual = FusedCrossEntropy.apply(logits, targets)
reference = F.cross_entropy(logits.float(), targets, reduction="sum")
upstream = torch.tensor(0.37, device=device)
dx = torch.autograd.grad(actual, logits, upstream)[0]
dxr = torch.autograd.grad(reference, logits, upstream)[0]
torch.testing.assert_close(actual, reference, rtol=2e-6, atol=2e-5)
torch.testing.assert_close(dx.float(), dxr.float(),
rtol=0.008 if dtype == torch.bfloat16 else 1e-4,
atol=2e-7)
def test_kernels(timeout_seconds=None):
if timeout_seconds is None:
timeout_seconds = KERNEL_TEST_TIMEOUT_SECONDS
deadline = time.monotonic() + timeout_seconds
try:
test_optimized_kernels(deadline=deadline)
for dtype in (torch.float32, torch.bfloat16):
_check_preflight_deadline(deadline, "Triton kernel preflight")
tol = 0.04 if dtype == torch.bfloat16 else 3e-4
x = torch.randn(
8, 2048, device="cuda", dtype=dtype, requires_grad=True
)
g = torch.randn_like(x)
y = TritonRMSFunction.apply(x, 1e-6)
dx = torch.autograd.grad(y, x, g)[0]
xr = x.detach().float().requires_grad_(True)
yr = xr * torch.rsqrt(xr.square().mean(-1, keepdim=True) + 1e-6)
dxr = torch.autograd.grad(yr, xr, g.float())[0]
torch.testing.assert_close(y.float(), yr, rtol=tol, atol=tol)
torch.testing.assert_close(dx.float(), dxr, rtol=tol, atol=tol)
a = torch.randn(
4096, device="cuda", dtype=dtype, requires_grad=True
)
b = torch.randn_like(a, requires_grad=True)
g = torch.randn_like(a)
z = TritonSwiGLUFunction.apply(a, b)
da, db = torch.autograd.grad(z, (a, b), g)
ar = a.detach().float().requires_grad_(True)
br = b.detach().float().requires_grad_(True)
zr = F.silu(ar) * br
dar, dbr = torch.autograd.grad(zr, (ar, br), g.float())
torch.testing.assert_close(z.float(), zr, rtol=tol, atol=tol)
torch.testing.assert_close(da.float(), dar, rtol=tol, atol=tol)
torch.testing.assert_close(db.float(), dbr, rtol=tol, atol=tol)
if USE_FUSED_MOE_COMBINE:
for d in (127, 2048):
_check_preflight_deadline(deadline, "Triton kernel preflight")
n, k = 129, 2
grouped = torch.randn(
n * k, d, device="cuda",
dtype=dtype, requires_grad=True,
)
weights = torch.softmax(
torch.randn(n, k, device="cuda"), dim=-1
).detach().requires_grad_(True)
inverse = torch.randperm(n * k, device="cuda")
upstream = torch.randn(n, d, device="cuda")
actual = FusedMoECombine.apply(grouped, weights, inverse)
dg, dw = torch.autograd.grad(
actual, (grouped, weights), upstream
)
gr = grouped.detach().clone().requires_grad_(True)
wr = weights.detach().clone().requires_grad_(True)
reference = combine_reference(gr, wr, inverse)
dgr, dwr = torch.autograd.grad(
reference, (gr, wr), upstream
)
torch.testing.assert_close(
actual, reference, rtol=1e-5, atol=1e-5
)
torch.testing.assert_close(
dg.float(), dgr.float(), rtol=1e-5, atol=1e-5
)
torch.testing.assert_close(
dw, dwr, rtol=3e-4, atol=3e-4
)
_check_preflight_deadline(deadline, "Triton kernel preflight")
torch.cuda.synchronize()
_check_preflight_deadline(deadline, "Triton kernel preflight")
except Exception as exc:
raise RuntimeError(
"Triton kernel preflight failed on the active GPU. Check the "
"CUDA/Triton/PyTorch versions and the reported assertion or timeout."
) from exc
print("All enabled Triton forward/backward checks passed.", flush=True)
# ===========================================================================
# ORION FLAGSHIP 2B — T2.2 ARCHITECTURE
# ===========================================================================
ARCH = {
"model_name": MODEL_NAME,
"implementation_version": 5,
"architecture": "orion_flagship_2b_t2_2",
"d_model": 2112,
"n_layers": 24,
"n_heads": 24,
"n_kv_heads": 6,
"rope_theta": 10000.0,
"sequence_length": SEQUENCE_LENGTH,
"tokenizer_name": TOKENIZER_NAME,
# Six reusable banks, each shared across four logical stages.
"n_procedure_banks": 6,
"procedure_bank_span": 4,
"n_experts": 12,
"top_k": 1,
"expert_hidden": 3456,
"shared_expert_hidden": 1280,
# Product-key Vault v2. Smaller capacity funds wider attention and experts;
# retain the state/stage-aware, confidence-gated read path.
"vault_key_parts": 128,
"vault_slots": 16_384,
"vault_top_component": 12,
"vault_top_k": 4,
"vault_gate_groups": 64,
"vault_read_layers": [3, 7, 11, 15, 19, 23],
"working_state_layers": [3, 7, 11, 15, 19, 23],
# One extra recurrent pass over the final six stages.
"deliberation_start": 18,
"deliberation_cycles": 1,
}
class Attention(nn.Module):
def __init__(self, cfg):
super().__init__()
d = cfg["d_model"]
self.n_heads = cfg["n_heads"]
self.n_kv_heads = cfg["n_kv_heads"]
if d % self.n_heads:
raise ValueError("d_model must be divisible by n_heads")
if self.n_heads % self.n_kv_heads:
raise ValueError("n_heads must be divisible by n_kv_heads")
self.head_dim = d // self.n_heads
if self.head_dim % 2:
raise ValueError("RoPE requires an even head dimension")
self.rope_theta = cfg["rope_theta"]
self.max_sequence_length = cfg["sequence_length"]
self.q = nn.Linear(d, d, bias=False)
self.k = nn.Linear(d, self.n_kv_heads * self.head_dim, bias=False)
self.v = nn.Linear(d, self.n_kv_heads * self.head_dim, bias=False)
self.o = nn.Linear(d, d, bias=False)
self.register_buffer("rope_cos", torch.empty(0), persistent=False)
self.register_buffer("rope_sin", torch.empty(0), persistent=False)
self.reset_rope()
def reset_rope(self):
device = self.q.weight.device
inv = 1.0 / (
self.rope_theta ** (
torch.arange(
0, self.head_dim, 2,
device=device, dtype=torch.float32,
) / self.head_dim
)
)
positions = torch.arange(
self.max_sequence_length,
device=device,
dtype=torch.float32,
)
angles = torch.outer(positions, inv)
self.rope_cos = angles.cos()
self.rope_sin = angles.sin()
def rope(self, x):
t = x.shape[-2]
cos = self.rope_cos[:t].to(x.dtype)[None, None]
sin = self.rope_sin[:t].to(x.dtype)[None, None]
even, odd = x[..., 0::2], x[..., 1::2]
return torch.stack(
(even * cos - odd * sin, even * sin + odd * cos),
dim=-1,
).flatten(-2)
def attention_gqa(self, x):
b, t, d = x.shape
q = self.q(x).view(b, t, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k(x).view(b, t, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v(x).view(b, t, self.n_kv_heads, self.head_dim).transpose(1, 2)
q, k = self.rope(q), self.rope(k)
y = F.scaled_dot_product_attention(
q, k, v,
is_causal=True,
dropout_p=0.0,
enable_gqa=True,
)
return self.o(y.transpose(1, 2).contiguous().view(b, t, d))
def attention_expand(self, x):
b, t, d = x.shape
q = self.q(x).view(b, t, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k(x).view(b, t, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v(x).view(b, t, self.n_kv_heads, self.head_dim).transpose(1, 2)
q, k = self.rope(q), self.rope(k)
repeats = self.n_heads // self.n_kv_heads
k = k.repeat_interleave(repeats, dim=1)
v = v.repeat_interleave(repeats, dim=1)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=0.0)
return self.o(y.transpose(1, 2).contiguous().view(b, t, d))
Attention.forward = attention_gqa
def select_attention(cfg, timeout_seconds=None):
if timeout_seconds is None:
timeout_seconds = ATTENTION_PROBE_TIMEOUT_SECONDS
hd = cfg["d_model"] // cfg["n_heads"]
deadline = time.monotonic() + timeout_seconds
def probe(native):
_check_preflight_deadline(deadline, "Flash attention preflight")
probe_tokens = max(1, min(256, cfg["sequence_length"]))
q = torch.randn(1, cfg["n_heads"], probe_tokens, hd,
device="cuda", dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(1, cfg["n_kv_heads"], probe_tokens, hd,
device="cuda", dtype=torch.bfloat16, requires_grad=True)
v = torch.randn_like(k, requires_grad=True)
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
if native:
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, enable_gqa=True
)
else:
repeats = cfg["n_heads"] // cfg["n_kv_heads"]
out = F.scaled_dot_product_attention(
q,
k.repeat_interleave(repeats, dim=1),
v.repeat_interleave(repeats, dim=1),
is_causal=True,
)
out.float().square().mean().backward()
torch.cuda.synchronize()
_check_preflight_deadline(deadline, "Flash attention preflight")
try:
probe(True)
eager = attention_gqa
print("Native Flash GQA forward/backward supported.")
except Exception as native_exc:
print("Native GQA probe failed; trying explicit KV expansion:", native_exc)
try:
probe(False)
except Exception as fallback_exc:
raise RuntimeError(
"Flash attention preflight failed for both native GQA and "
"explicit KV expansion. Check GPU capability, BF16 support, "
"and the PyTorch CUDA attention backend."
) from fallback_exc
eager = attention_expand
print("Using explicit KV expansion with Flash attention.")
Attention.forward = eager
if not COMPILE_ATTENTION:
return
trial = x = out = None
try:
_check_preflight_deadline(deadline, "Attention compilation preflight")
with torch.device("cuda"):
trial = Attention(cfg)
compiled = torch.compile(eager, fullgraph=True, dynamic=False)
probe_batch = max(1, min(MICRO_BATCH_SIZE, 2))
probe_tokens = max(1, min(SEQUENCE_LENGTH, 256))
x = torch.randn(
probe_batch, probe_tokens, cfg["d_model"],
device="cuda", dtype=torch.float32, requires_grad=True,
)
with torch.autocast("cuda", dtype=torch.bfloat16,
cache_enabled=AUTOCAST_CACHE):
out = compiled(trial, x)
out.float().square().mean().backward()
torch.cuda.synchronize()
_check_preflight_deadline(deadline, "Attention compilation preflight")
Attention.forward = compiled
print("Compiled attention startup forward/backward passed.")
except TimeoutError as exc:
raise RuntimeError(
"Attention compilation preflight exceeded its time limit; "
"inspect the CUDA/PyTorch compiler setup before training."
) from exc
except Exception as exc:
Attention.forward = eager
print("Attention compilation failed; using eager attention:", exc)
finally:
del trial, x, out
gc.collect()
torch.cuda.empty_cache()
class Expert(nn.Module):
def __init__(self, d, hidden):
super().__init__()
self.gate = nn.Linear(d, hidden, bias=False)
self.up = nn.Linear(d, hidden, bias=False)
self.down = nn.Linear(hidden, d, bias=False)
def forward(self, x):
return self.down(TritonSwiGLUFunction.apply(self.gate(x), self.up(x)))
class ProcedureBank(nn.Module):
"""T2.1 reusable transformations with state/stage-aware top-1 routing."""
def __init__(self, cfg):
super().__init__()
d = cfg["d_model"]
self.n_experts = cfg["n_experts"]
self.top_k = cfg["top_k"]
if self.top_k != 1:
raise ValueError("T2.1 ProcedureBank expects top_k=1")
self.expert_hidden = cfg["expert_hidden"]
self.shared_hidden = cfg["shared_expert_hidden"]
# Final logit is the soft NULL/no-specialist path.
self.router = nn.Linear(d, self.n_experts + 1, bias=True)
self.experts = nn.ModuleList([
Expert(d, self.expert_hidden) for _ in range(self.n_experts)
])
self.shared = Expert(d, self.shared_hidden)
self.shared_gate = nn.Linear(d, 1, bias=True)
self.specialization_scale = 0.0
self.telemetry_enabled = False
self.reset_telemetry()
def set_specialization_scale(self, value):
self.specialization_scale = float(max(0.0, min(1.0, value)))
def reset_telemetry(self):
self._telemetry = {
"router_entropy_sum": 0.0,
"load_entropy_sum": 0.0,
"max_load_sum": 0.0,
"min_load_sum": 0.0,
"dead_experts_sum": 0.0,
"null_probability_sum": 0.0,
"shared_gate_sum": 0.0,
"specialist_gate_sum": 0.0,
"calls": 0,
}
def forward(self, x, routing_context):
shape = x.shape
flat = x.reshape(-1, shape[-1])
route_flat = routing_context.reshape(-1, shape[-1])
n = flat.shape[0]
e = self.n_experts
k = 1
with torch.autocast("cuda", enabled=False):
logits = F.linear(
route_flat.float(), self.router.weight.float(), self.router.bias.float()
)
probabilities = F.softmax(logits, dim=-1)
p_null = probabilities[:, -1]
nonnull = probabilities[:, :-1]
nonnull_norm = nonnull / nonnull.sum(-1, keepdim=True).clamp_min(1e-9)
selected_p, selected_e = torch.max(nonnull_norm, dim=-1)
assignments = selected_e
counts_gpu = torch.zeros(e, device=flat.device, dtype=torch.int32)
counts_gpu.scatter_add_(
0, assignments,
torch.ones_like(assignments, dtype=torch.int32),
)
load = counts_gpu.float() / max(1, n)
balance = e * (nonnull_norm.mean(0) * load).sum()
z_loss = logits.logsumexp(-1).square().mean()
# Nonnegative specialization objective:
# low per-token entropy + high marginal entropy.
log_e = math.log(float(e))
token_h = -(
nonnull_norm * nonnull_norm.clamp_min(1e-9).log()
).sum(-1).mean()
marginal = nonnull_norm.mean(0)
marginal_h = -(
marginal * marginal.clamp_min(1e-9).log()
).sum()
specialization = token_h / log_e + (1.0 - marginal_h / log_e)
# Preserve forward specialist scale near 1 at initialization while
# allowing a bounded task gradient into the selected router score.
ratio = selected_p / selected_p.detach().clamp_min(1e-3)
task_st = 1.0 + ROUTER_TASK_GRAD_SCALE * (ratio - 1.0)
specialist_gate = (1.0 - p_null) * task_st
weights = specialist_gate.unsqueeze(-1)
order = torch.argsort(assignments)
inverse = torch.empty_like(order)
inverse.scatter_(
0, order,
torch.arange(order.numel(), device=order.device),
)
compute_dtype = (
torch.get_autocast_dtype("cuda")
if torch.is_autocast_enabled("cuda") else flat.dtype
)
packed_input = FusedMoEPack.apply(
flat, order, inverse, k, compute_dtype
) if USE_FUSED_MOE_PACK else flat.to(compute_dtype).index_select(0, order)
counts = counts_gpu.cpu().tolist()
chunks = packed_input.split(counts, dim=0)
pieces = []
for expert, chunk, count in zip(self.experts, chunks, counts):
if count:
pieces.append(expert(chunk))
if not pieces:
raise RuntimeError("Procedure router produced no specialist assignments.")
routed = torch.cat(pieces, dim=0)
if USE_FUSED_MOE_COMBINE:
routed = FusedMoECombine.apply(routed, weights, inverse)
else:
routed = combine_reference(routed, weights, inverse)
shared = self.shared(flat.to(compute_dtype)).float()
shared_mix = torch.sigmoid(self.shared_gate(flat.float()))
output = routed + shared_mix * shared
if self.telemetry_enabled:
router_entropy = -(
nonnull_norm * nonnull_norm.clamp_min(1e-9).log()
).sum(-1).mean()
load_entropy = -(
load * load.clamp_min(1e-9).log()
).sum()
self._telemetry["router_entropy_sum"] += float(router_entropy)
self._telemetry["load_entropy_sum"] += float(load_entropy)
self._telemetry["max_load_sum"] += float(load.max())
self._telemetry["min_load_sum"] += float(load.min())
self._telemetry["dead_experts_sum"] += float((counts_gpu == 0).sum())
self._telemetry["null_probability_sum"] += float(p_null.mean())
self._telemetry["shared_gate_sum"] += float(shared_mix.mean())
self._telemetry["specialist_gate_sum"] += float(specialist_gate.mean())
self._telemetry["calls"] += 1
auxiliary = (
ROUTER_AUX_COEF * balance
+ ROUTER_Z_COEF * z_loss
+ ROUTER_SPECIALIZATION_COEF * self.specialization_scale * specialization
)
return output.view(shape), auxiliary
class KnowledgeVault(nn.Module):
"""T2.1 confidence-gated, state/stage-aware product-key memory."""
def __init__(self, cfg):
super().__init__()
d = cfg["d_model"]
parts = cfg["vault_key_parts"]
slots = cfg["vault_slots"]
groups = cfg["vault_gate_groups"]
if parts * parts != slots:
raise ValueError("vault_slots must equal vault_key_parts**2")
if d % 2:
raise ValueError("d_model must be even for split product keys")
if d % groups:
raise ValueError("d_model must be divisible by vault_gate_groups")
self.d = d
self.parts = parts
self.component_top = cfg["vault_top_component"]
self.top_k = cfg["vault_top_k"]
self.groups = groups
self.group_width = d // groups
self.slots = slots
self.query = nn.Linear(d, d, bias=False)
self.state_query = nn.Linear(d, d, bias=False)
self.key_a = nn.Parameter(torch.empty(parts, d // 2))
self.key_b = nn.Parameter(torch.empty(parts, d // 2))
self.values = nn.Parameter(torch.empty(slots, d))
# Learned residual delta from [token, retrieved memory].
self.fusion = nn.Linear(2 * d, d, bias=False)
# Three bounded confidence features + token representation -> grouped gate.
self.gate = nn.Linear(d + 3, groups, bias=True)
self.telemetry_enabled = False
self.reset_telemetry()
def reset_telemetry(self):
self._telemetry = {
"gate_sum": 0.0,
"gate_std_sum": 0.0,
"retrieval_entropy_sum": 0.0,
"top1_weight_sum": 0.0,
"margin_conf_sum": 0.0,
"unique_slots_sum": 0.0,
"calls": 0,
}
def forward(self, x, state, stage_context):
shape = x.shape
flat = x.reshape(-1, self.d)
state_flat = state.reshape(-1, self.d)
context = stage_context.reshape(1, self.d).expand_as(flat)
q_input = flat + self.state_query(state_flat) + context
q = self.query(q_input)
qa, qb = q.chunk(2, dim=-1)
# Scale product-key scores to keep confidence numerically well behaved.
score_scale = 1.0 / math.sqrt(self.d / 2)
score_a = F.linear(qa.float(), self.key_a.float()) * score_scale
score_b = F.linear(qb.float(), self.key_b.float()) * score_scale
sa, ia = torch.topk(score_a, self.component_top, dim=-1)
sb, ib = torch.topk(score_b, self.component_top, dim=-1)
candidate_scores = (sa.unsqueeze(2) + sb.unsqueeze(1)).flatten(1)
candidate_ids = (
ia.unsqueeze(2) * self.parts + ib.unsqueeze(1)
).flatten(1)
best_scores, best_pos = torch.topk(
candidate_scores, self.top_k, dim=-1
)
best_ids = candidate_ids.gather(1, best_pos)
weights_f = F.softmax(best_scores, dim=-1)
gathered = self.values.index_select(0, best_ids.reshape(-1))
gathered = gathered.view(flat.shape[0], self.top_k, self.d)
memory = (gathered * weights_f.to(gathered.dtype).unsqueeze(-1)).sum(dim=1)
compute_dtype = (
torch.get_autocast_dtype("cuda")
if torch.is_autocast_enabled("cuda") else flat.dtype
)
delta = self.fusion(
torch.cat((flat.to(compute_dtype), memory.to(compute_dtype)), dim=-1)
)
entropy = -(
weights_f * weights_f.clamp_min(1e-9).log()
).sum(-1)
entropy_conf = 1.0 - entropy / math.log(float(self.top_k))
top1 = weights_f[:, 0]
if self.top_k > 1:
margin_conf = torch.sigmoid(best_scores[:, 0] - best_scores[:, 1])
else:
margin_conf = torch.ones_like(top1)
confidence = torch.stack((top1, margin_conf, entropy_conf), dim=-1)
with torch.autocast("cuda", enabled=False):
gate_features = torch.cat((flat.float(), confidence.float()), dim=-1)
gate_logits = F.linear(
gate_features, self.gate.weight.float(), self.gate.bias.float()
)
group_gate = torch.sigmoid(gate_logits)
gate = group_gate.unsqueeze(-1).expand(
-1, self.groups, self.group_width
).reshape(-1, self.d).to(delta.dtype)
if self.telemetry_enabled:
unique_slots = torch.unique(best_ids).numel()
self._telemetry["gate_sum"] += float(gate.float().mean())
self._telemetry["gate_std_sum"] += float(gate.float().std())
self._telemetry["retrieval_entropy_sum"] += float(entropy.mean())
self._telemetry["top1_weight_sum"] += float(top1.mean())
self._telemetry["margin_conf_sum"] += float(margin_conf.mean())
self._telemetry["unique_slots_sum"] += float(unique_slots)
self._telemetry["calls"] += 1
return (gate * delta).view(shape)
class WorkingState(nn.Module):
"""Bounded causal residual state retained from the stable T2 design."""
def __init__(self, d):
super().__init__()
self.read = nn.Linear(d, d, bias=False)
self.read_gate = nn.Linear(d, 1, bias=True)
self.write_gate = nn.Linear(2 * d, d, bias=True)
self.write_value = nn.Linear(2 * d, d, bias=False)
self.telemetry_enabled = False
self.reset_telemetry()
def reset_telemetry(self):
self._telemetry = {
"read_gate_sum": 0.0,
"read_calls": 0,
"write_gate_sum": 0.0,
"state_delta_rms_sum": 0.0,
"state_rms_sum": 0.0,
"write_calls": 0,
}
def inject(self, x, state):
gate = torch.sigmoid(self.read_gate(x.float())).to(x.dtype)
if self.telemetry_enabled:
self._telemetry["read_gate_sum"] += float(gate.float().mean())
self._telemetry["read_calls"] += 1
return x + gate * self.read(state)
def update(self, x, state):
pair = torch.cat((x, state), dim=-1)
gate = torch.sigmoid(self.write_gate(pair))
value = torch.tanh(self.write_value(pair))
# Critical stability property: convex replacement, never unbounded add.
updated = (1.0 - gate) * state + gate * value
if self.telemetry_enabled:
delta_rms = (updated.float() - state.float()).square().mean().sqrt()
state_rms = updated.float().square().mean().sqrt()
self._telemetry["write_gate_sum"] += float(gate.float().mean())
self._telemetry["state_delta_rms_sum"] += float(delta_rms)
self._telemetry["state_rms_sum"] += float(state_rms)
self._telemetry["write_calls"] += 1
return updated
class CortexStage(nn.Module):
def __init__(self, cfg):
super().__init__()
self.norm1 = RMSNorm()
self.attn = Attention(cfg)
self.norm2 = RMSNorm()
def forward(self, x, bank, routing_add, disable_procedure=False):
x = x + self.attn(self.norm1(x))
if disable_procedure:
return x, x.new_zeros(())
procedural_input = self.norm2(x)
routing_context = procedural_input + routing_add
y, auxiliary = bank(procedural_input, routing_context)
return x + y, auxiliary
class OrionFlagshipT21(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = dict(cfg)
d = cfg["d_model"]
self.embedding = nn.Embedding(cfg["vocab_size"], d)
self.blocks = nn.ModuleList([
CortexStage(cfg) for _ in range(cfg["n_layers"])
])
self.procedure_banks = nn.ModuleList([
ProcedureBank(cfg) for _ in range(cfg["n_procedure_banks"])
])
self.vault = KnowledgeVault(cfg)
self.working = WorkingState(d)
self.state_seed = nn.Parameter(torch.zeros(d))
# Shared context signals for routing. They start at zero so T2.1 begins
# close to the known-stable T2 information flow and earns the additions.
self.route_state_proj = nn.Linear(d, d, bias=False)
self.stage_embedding = nn.Embedding(cfg["n_layers"], d)
self.cycle_embedding = nn.Embedding(cfg["deliberation_cycles"] + 1, d)
self.norm = RMSNorm()
self.vault_read_layers = set(cfg["vault_read_layers"])
self.working_state_layers = set(cfg["working_state_layers"])
self.bank_span = cfg["procedure_bank_span"]
self.deliberation_start = cfg["deliberation_start"]
self.deliberation_cycles = cfg["deliberation_cycles"]
def bank_for_layer(self, layer_index):
return self.procedure_banks[layer_index // self.bank_span]
def stage_context(self, layer_index, cycle_index, dtype):
return (
self.stage_embedding.weight[layer_index]
+ self.cycle_embedding.weight[cycle_index]
).to(dtype)
def set_specialization_scale(self, value):
for bank in self.procedure_banks:
bank.set_specialization_scale(value)
def set_telemetry(self, enabled):
self.vault.telemetry_enabled = enabled
self.working.telemetry_enabled = enabled
for bank in self.procedure_banks:
bank.telemetry_enabled = enabled
def reset_telemetry(self):
self.vault.reset_telemetry()
self.working.reset_telemetry()
for bank in self.procedure_banks:
bank.reset_telemetry()
def telemetry_summary(self):
def avg(total, count):
return total / max(1, count)
v = self.vault._telemetry
w = self.working._telemetry
banks = []
for index, bank in enumerate(self.procedure_banks):
t = bank._telemetry
calls = max(1, t["calls"])
banks.append({
"bank": index,
"router_entropy": t["router_entropy_sum"] / calls,
"load_entropy": t["load_entropy_sum"] / calls,
"max_load": t["max_load_sum"] / calls,
"min_load": t["min_load_sum"] / calls,
"dead_experts": t["dead_experts_sum"] / calls,
"null_probability": t["null_probability_sum"] / calls,
"shared_gate": t["shared_gate_sum"] / calls,
"specialist_gate": t["specialist_gate_sum"] / calls,
})
vcalls = max(1, v["calls"])
return {
"vault_gate": v["gate_sum"] / vcalls,
"vault_gate_std": v["gate_std_sum"] / vcalls,
"vault_entropy": v["retrieval_entropy_sum"] / vcalls,
"vault_top1": v["top1_weight_sum"] / vcalls,
"vault_margin": v["margin_conf_sum"] / vcalls,
"vault_unique": v["unique_slots_sum"] / vcalls,
"state_read": avg(w["read_gate_sum"], w["read_calls"]),
"state_write": avg(w["write_gate_sum"], w["write_calls"]),
"state_delta": avg(w["state_delta_rms_sum"], w["write_calls"]),
"state_rms": avg(w["state_rms_sum"], w["write_calls"]),
"banks": banks,
}
def initialize_weights(self):
for module in self.modules():
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if getattr(module, "bias", None) is not None:
nn.init.zeros_(module.bias)
nn.init.normal_(self.vault.key_a, std=0.02)
nn.init.normal_(self.vault.key_b, std=0.02)
nn.init.normal_(self.vault.values, std=0.02)
nn.init.zeros_(self.state_seed)
# New context pathways begin neutral/closed for stability.
nn.init.zeros_(self.route_state_proj.weight)
nn.init.zeros_(self.stage_embedding.weight)
nn.init.zeros_(self.cycle_embedding.weight)
nn.init.zeros_(self.vault.state_query.weight)
nn.init.zeros_(self.vault.gate.weight)
nn.init.constant_(self.vault.gate.bias, -2.5)
nn.init.constant_(self.working.read_gate.bias, -2.0)
for bank in self.procedure_banks:
nn.init.constant_(bank.shared_gate.bias, -1.0)
nn.init.zeros_(bank.router.bias)
# Soft NULL is available but initially disfavoured.
with torch.no_grad():
bank.router.bias[-1] = -2.0
residual_std = 0.02 / math.sqrt(2 * self.cfg["n_layers"])
for block in self.blocks:
nn.init.normal_(block.attn.o.weight, std=residual_std)
for bank in self.procedure_banks:
for expert in list(bank.experts) + [bank.shared]:
nn.init.normal_(expert.down.weight, std=residual_std)
nn.init.normal_(self.vault.fusion.weight, std=residual_std)
def loss_chunk(self, hidden, targets):
logits = F.linear(hidden, self.embedding.weight)
if USE_FUSED_CROSS_ENTROPY and logits.shape[-1] <= 65536:
return FusedCrossEntropy.apply(logits, targets)
return F.cross_entropy(logits, targets, reduction="sum")
def run_stage(self, index, x, state, cycle_index=0, ablate=frozenset()):
disable_state = "state" in ablate
if not disable_state:
x = self.working.inject(x, state)
context = self.stage_context(index, cycle_index, x.dtype)
routing_add = context
if not disable_state:
routing_add = routing_add + self.route_state_proj(state)
bank = self.bank_for_layer(index)
x, auxiliary = self.blocks[index](
x,
bank,
routing_add,
disable_procedure=("procedure" in ablate),
)
if index in self.vault_read_layers and "vault" not in ablate:
x = x + self.vault(x, state, context)
if index in self.working_state_layers and not disable_state:
state = self.working.update(x, state)
return x, state, auxiliary
def forward(self, input_ids, labels, ablate=None):
ablate = frozenset() if ablate is None else frozenset(ablate)
x = self.embedding(input_ids)
state = self.state_seed.to(x.dtype).view(1, 1, -1).expand_as(x).clone()
auxiliary = x.new_zeros(())
for index in range(len(self.blocks)):
if self.training and index < max(0, len(self.blocks) - UNCHECKPOINTED_LAST_N):
x, state, aux = checkpoint(
lambda a, s, idx=index: self.run_stage(
idx, a, s, cycle_index=0, ablate=ablate
),
x, state,
use_reentrant=False,
preserve_rng_state=False,
)
else:
x, state, aux = self.run_stage(
index, x, state, cycle_index=0, ablate=ablate
)
auxiliary = auxiliary + aux
executed_stages = len(self.blocks)
if "deliberation" not in ablate:
for cycle in range(1, self.deliberation_cycles + 1):
for index in range(self.deliberation_start, len(self.blocks)):
if self.training and CHECKPOINT_DELIBERATION:
x, state, aux = checkpoint(
lambda a, s, idx=index, cyc=cycle: self.run_stage(
idx, a, s, cycle_index=cyc, ablate=ablate
),
x, state,
use_reentrant=False,
preserve_rng_state=False,
)
else:
x, state, aux = self.run_stage(
index, x, state, cycle_index=cycle, ablate=ablate
)
auxiliary = auxiliary + aux
executed_stages += 1
if "state" not in ablate:
x = self.working.inject(x, state)
x = self.norm(x)
hidden = x.reshape(-1, x.shape[-1])
targets = labels.reshape(-1)
if LOSS_CHUNK_TOKENS is None:
ce = self.loss_chunk(hidden, targets) / targets.numel()
else:
ce_sum = hidden.new_zeros((), dtype=torch.float32)
for start in range(0, hidden.shape[0], LOSS_CHUNK_TOKENS):
h = hidden[start:start + LOSS_CHUNK_TOKENS]
target = targets[start:start + LOSS_CHUNK_TOKENS]
if self.training and CHECKPOINT_LOSS_CHUNKS:
part = checkpoint(
self.loss_chunk, h, target,
use_reentrant=False,
preserve_rng_state=False,
)
else:
part = self.loss_chunk(h, target)
ce_sum = ce_sum + part
ce = ce_sum / targets.numel()
return ce, auxiliary / max(1, executed_stages)
@torch.no_grad()
def run_telemetry_probe(model, input_ids, labels):
"""All ranks must enter; forwards must go through the DDP wrapper."""
module = model.module if isinstance(model, DDP) else model
was_training = model.training
try:
model.eval()
module.reset_telemetry()
module.set_telemetry(True)
with torch.autocast(
"cuda", dtype=torch.bfloat16, cache_enabled=AUTOCAST_CACHE
):
ce, auxiliary = model(input_ids, labels)
stats = module.telemetry_summary()
stats["probe_ce"] = float(ce)
stats["probe_aux"] = float(auxiliary)
return stats
finally:
module.set_telemetry(False)
model.train(was_training)
@torch.no_grad()
def run_ablation_probe(model, input_ids, labels):
"""All ranks run each ablation through the distributed wrapper."""
was_training = model.training
try:
model.eval()
results = {}
with torch.autocast(
"cuda", dtype=torch.bfloat16, cache_enabled=AUTOCAST_CACHE
):
base, _ = model(input_ids, labels)
results["full"] = float(base)
for name in ("vault", "state", "procedure", "deliberation"):
ce, _ = model(input_ids, labels, ablate={name})
results[f"no_{name}"] = float(ce)
return results
finally:
model.train(was_training)
def parameter_count_breakdown(cfg):
"""Exact constructor algebra (tied embedding counted once), no allocation."""
d = cfg["d_model"]
layers = cfg["n_layers"]
e = cfg["n_experts"]
groups = cfg["vault_gate_groups"]
hd = d // cfg["n_heads"]
return {
"embedding": cfg["vocab_size"] * d,
"cortex": layers * (2 * d * d + 2 * d * cfg["n_kv_heads"] * hd),
"procedure": cfg["n_procedure_banks"] * (
3 * d * (e * cfg["expert_hidden"] + cfg["shared_expert_hidden"])
+ d * (e + 1) + (e + 1) + d + 1
),
"vault": 4 * d * d + cfg["vault_key_parts"] * d
+ cfg["vault_slots"] * d + groups * (d + 4),
"state": 5 * d * d + 2 * d + 1,
"context": d + d * d + layers * d + (cfg["deliberation_cycles"] + 1) * d,
}
def parameter_counts(cfg):
"""Exact meta count + approximate per-stage active compute proxy."""
with torch.device("meta"):
probe = OrionFlagshipT21(cfg)
total = sum(p.numel() for p in probe.parameters())
expected = sum(parameter_count_breakdown(cfg).values())
if total != expected:
raise RuntimeError(f"Parameter-count algebra disagrees with model: {expected} != {total}")
d = cfg["d_model"]
hd = d // cfg["n_heads"]
attn = 2 * d * d + 2 * d * cfg["n_kv_heads"] * hd
specialist = 3 * d * cfg["expert_hidden"]
shared = 3 * d * cfg["shared_expert_hidden"]
router = d * (cfg["n_experts"] + 1) + (cfg["n_experts"] + 1)
active = cfg["vocab_size"] * d
active += cfg["n_layers"] * (attn + router + specialist + shared + d * d)
# Vault/working paths are intentionally approximate compute-equivalents.
vault_dense = 4 * d * d + cfg["vault_top_k"] * d
active += len(cfg["vault_read_layers"]) * vault_dense
working_dense = 5 * d * d
active += len(cfg["working_state_layers"]) * working_dense
return total, active
def gradient_health_summary(model):
# DDP gradients are already replicated/averaged after synchronized backward.
# Summing squared norms across ranks would inflate them by sqrt(WORLD_SIZE).
groups = ("cortex", "procedure", "vault", "state", "embedding")
first_parameter = next(model.parameters(), None)
device = first_parameter.device if first_parameter is not None else DEVICE
squared_norms = torch.zeros(len(groups), device=device, dtype=torch.float32)
for name, p in model.named_parameters():
if p.grad is None:
continue
name = name.removeprefix("module.")
g = p.grad.detach().float()
value = g.square().sum()
if name.startswith("procedure_banks") or name.startswith("route_state_proj") \
or name.startswith("stage_embedding") or name.startswith("cycle_embedding"):
group = "procedure"
elif name.startswith("vault"):
group = "vault"
elif name.startswith("working") or name.startswith("state_seed"):
group = "state"
elif name.startswith("embedding"):
group = "embedding"
else:
group = "cortex"
squared_norms[groups.index(group)] += value
return dict(zip(groups, squared_norms.sqrt().tolist()))
def health_warnings(stats, grad_stats, step):
warnings = []
if not math.isfinite(stats.get("probe_ce", float("nan"))):
warnings.append("NONFINITE probe CE")
if step >= 10_000:
if stats["vault_gate"] < 0.005:
warnings.append("Vault gate is nearly closed")
if stats["vault_gate"] > 0.98:
warnings.append("Vault gate is saturated open")
if stats["state_delta"] < 1e-5:
warnings.append("Working State delta is nearly zero")
for bank in stats["banks"]:
if bank["dead_experts"] > 0:
warnings.append(f"B{bank['bank']} has dead experts in probe")
if bank["max_load"] > 0.50:
warnings.append(f"B{bank['bank']} routing load >50% on one expert")
if step >= 10_000 and bank["null_probability"] > 0.90:
warnings.append(f"B{bank['bank']} soft-NULL probability >90%")
if step >= 100:
for key in ("procedure", "vault", "state"):
if grad_stats.get(key, 0.0) < 1e-10:
warnings.append(f"{key} gradient norm is ~zero")
return warnings
def run_startup_model_probe(model, vocab_size):
"""One production-shape F/B pass before real data is consumed."""
if not STARTUP_MODEL_PROBE:
return
print("Running T2.1 production-shape startup health probe...", flush=True)
model.train()
model.zero_grad(set_to_none=True)
x = torch.randint(
0, vocab_size, (MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device="cuda", dtype=torch.long,
)
y = torch.randint(
0, vocab_size, (MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device="cuda", dtype=torch.long,
)
module = model.module if isinstance(model, DDP) else model
module.set_specialization_scale(0.0)
with torch.autocast("cuda", dtype=torch.bfloat16, cache_enabled=AUTOCAST_CACHE):
ce, aux = model(x, y)
loss = ce + aux
finite_loss = torch.isfinite(loss.detach()).to(dtype=torch.int32)
if dist.is_initialized():
dist.all_reduce(finite_loss, op=dist.ReduceOp.MIN)
if not bool(finite_loss.item()):
raise RuntimeError("Startup model probe produced nonfinite loss.")
loss.backward()
grads = gradient_health_summary(model)
for key in ("cortex", "procedure", "vault", "state", "embedding"):
if not math.isfinite(grads[key]) or grads[key] <= 0.0:
raise RuntimeError(f"Startup probe: invalid {key} gradient norm {grads[key]}")
model.zero_grad(set_to_none=True)
stats = run_telemetry_probe(model, x, y)
print(
f"Startup probe PASS: ce={float(ce):.4f} aux={float(aux):.4f} "
f"grads={grads} vault_gate={stats['vault_gate']:.3f} "
f"state_delta={stats['state_delta']:.4f}",
flush=True,
)
del x, y, ce, aux, loss
gc.collect()
torch.cuda.empty_cache()
# ===========================================================================
# Immutable data specification
# ===========================================================================
def create_data_spec(vocab_size):
api = HfApi(token=HF_TOKEN)
info = api.repo_info(
repo_id=DATA_REPO_ID,
repo_type=DATA_REPO_TYPE,
revision=DATA_REVISION,
)
revision = info.sha
path = hf_hub_download(
repo_id=DATA_REPO_ID,
repo_type=DATA_REPO_TYPE,
revision=revision,
filename=DATA_MANIFEST,
token=HF_TOKEN,
cache_dir=str(WORK / "shard_cache"),
)
manifest = json.loads(Path(path).read_text())
if manifest.get("format_version") != 3:
raise ValueError(
f"Expected Nano-base manifest format_version=3, got "
f"{manifest.get('format_version')!r}."
)
if manifest.get("phase") != "orion-nano-base-v1":
raise ValueError(
f"Wrong data phase in manifest: {manifest.get('phase')!r}."
)
target_tokens = int(manifest.get("target_tokens", 0))
written_tokens = int(manifest.get("total_tokens_written", 0))
if target_tokens != DATA_TARGET_TOKENS or written_tokens != target_tokens:
raise ValueError(
"Nano base corpus is not complete yet: "
f"written={written_tokens:,}, target={target_tokens:,}. "
"Wait for the pretokenizer to finish before starting training."
)
expected_targets = {
"general_fineweb_edu": 30_000_000_000,
"math_finemath_4plus": 6_000_000_000,
"math_openwebmath": 4_000_000_000,
"code_python_clean": 10_000_000_000,
}
if manifest.get("source_targets") != expected_targets:
raise ValueError(
"Nano corpus source budgets differ from the frozen 60/20/20 mix."
)
dtypes = {
"uint16": "<u2",
"uint32": "<u4",
"<u2": "<u2",
"<u4": "<u4",
}
dtype = dtypes.get(manifest.get("dtype"))
if dtype is None:
raise ValueError(f"Unsupported token dtype: {manifest.get('dtype')}")
if dtype == "<u2" and vocab_size > 65536:
raise ValueError("uint16 data cannot represent this vocabulary.")
tokenizer_name = manifest.get("tokenizer_name")
if tokenizer_name != TOKENIZER_NAME:
raise ValueError(
f"Data tokenizer {tokenizer_name!r} differs from {TOKENIZER_NAME!r}."
)
if int(manifest.get("vocab_size", -1)) != vocab_size:
raise ValueError("Data manifest vocabulary differs from model tokenizer.")
files = []
for worker in manifest.get("workers", {}).values():
for filename in worker.get("shard_files", []):
if not isinstance(filename, str):
raise ValueError("Expected filenames in manifest shard_files.")
filename = safe_relative_path(filename)
prefix = DATA_CACHE_DIR.rstrip("/") + "/"
if not filename.startswith(prefix):
filename = prefix + filename
files.append(filename)
files = sorted(set(files))
if not files:
raise ValueError("No Nano token shards are available in the manifest.")
print(
f"Nano dataset: {len(files)} shards pinned to revision {revision}.\n"
f"Manifest tokens: {written_tokens:,}; train budget: {TRAIN_TOKENS:,}.\n"
"Shard lengths are read from actual files; v3 training interleaves sources every sequence.",
flush=True,
)
return {
"version": 2,
"kind": "nano_base_hub",
"repo_id": DATA_REPO_ID,
"repo_type": DATA_REPO_TYPE,
"revision": revision,
"dtype": dtype,
"files": files,
"vocab_size": vocab_size,
"manifest_target_tokens": target_tokens,
"source_targets": expected_targets,
}
# ===========================================================================
# Stratified multi-shard data stream (v3)
# ===========================================================================
class ShardStream:
"""
Deterministic source-stratified interleaving across the whole Nano corpus.
Why this exists:
v2 shuffled shard order, but then consumed one entire ~268M-token shard
before moving on. Because shards are source-homogeneous, that created very
long single-domain runs (for example, hundreds of millions of Python
tokens in a row).
v3 fixes that without rewriting the 50B-token cache:
* Source selection follows an exact 25-sequence cycle:
15 FineWeb-Edu, 3 FineMath-4+, 2 OpenWebMath, 5 Code.
* The 25 source labels are independently shuffled every cycle.
* Within each source, shards are shuffled per source-epoch.
* Within each shard, 1024-token sequences are shuffled.
* Up to four source shards are naturally active at once, with an LRU
memmap cache and asynchronous next-shard downloads.
* The complete per-source cursor + global sequence index is checkpointed,
making resume exact and independent of batch boundaries.
No training sequence crosses a shard boundary. Tiny shard tails are omitted.
"""
SOURCE_NAMES = (
"general_fineweb_edu",
"math_finemath_4plus",
"math_openwebmath",
"code_python_clean",
)
def __init__(self, spec, seed, cursor=None):
self.spec = copy.deepcopy(spec)
self.seed = int(seed)
self.files = list(spec["files"])
self.dtype = np.dtype(spec["dtype"])
if spec.get("version") != 2 or spec.get("kind") != "nano_base_hub":
raise ValueError("This script expects a version-2 Nano-base Hub data spec.")
if not self.files:
raise ValueError("Empty data file list.")
# Split the immutable file list by the source id embedded by the
# pretokenizer in each shard filename.
self.source_files = {name: [] for name in self.SOURCE_NAMES}
for filename in self.files:
matches = [name for name in self.SOURCE_NAMES if f"-{name}.bin" in filename]
if len(matches) != 1:
raise ValueError(
f"Could not uniquely identify Nano source for shard: {filename}"
)
self.source_files[matches[0]].append(filename)
for source, names in self.source_files.items():
if not names:
raise ValueError(f"No shards found for required source {source!r}.")
names.sort()
# global_sequence chooses the source deterministically. Each source has
# an independent cursor through its own shuffled shards/sequences.
self.cursor = {
"stream_version": DATA_STREAM_VERSION,
"global_sequence": 0,
"sources": {
source: {
"epoch": 0,
"shard_position": 0,
"sequence_position": 0,
}
for source in self.SOURCE_NAMES
},
}
if cursor is not None:
if int(cursor.get("stream_version", -1)) != DATA_STREAM_VERSION:
raise ValueError(
f"Saved data stream version {cursor.get('stream_version')!r} "
f"is incompatible with v{DATA_STREAM_VERSION} stratified mixing."
)
self.cursor = copy.deepcopy(cursor)
if int(self.cursor["global_sequence"]) < 0:
raise ValueError("Negative global data cursor.")
for source in self.SOURCE_NAMES:
c = self.cursor["sources"].get(source)
if not isinstance(c, dict):
raise ValueError(f"Missing saved cursor for source {source!r}.")
if any(int(c[k]) < 0 for k in ("epoch", "shard_position", "sequence_position")):
raise ValueError(f"Negative data cursor for source {source!r}.")
if int(c["shard_position"]) >= len(self.source_files[source]):
raise ValueError(f"Invalid shard position for source {source!r}.")
self.cache_dir = WORK / "shard_cache"
self.cache_dir.mkdir(parents=True, exist_ok=True)
self._mapped = OrderedDict()
self._active = {}
self._shard_orders = {}
self._sequence_orders = {}
self._download_pool = ThreadPoolExecutor(
max_workers=DOWNLOAD_WORKERS, thread_name_prefix="shard-download"
)
self._futures = {}
def state_dict(self):
return copy.deepcopy(self.cursor)
def _source_cycle(self, cycle):
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, int(cycle), 303])
)
order = np.asarray(MIX_CYCLE, dtype=object)
return order[rng.permutation(len(order))]
def _source_for_global(self, global_sequence):
cycle, offset = divmod(int(global_sequence), len(MIX_CYCLE))
return str(self._source_cycle(cycle)[offset])
def _shard_order(self, source, epoch):
key = (source, int(epoch))
if key not in self._shard_orders:
source_id = self.SOURCE_NAMES.index(source)
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, source_id, int(epoch), 101])
)
self._shard_orders[key] = rng.permutation(len(self.source_files[source]))
return self._shard_orders[key]
def _download(self, name):
return hf_hub_download(
repo_id=self.spec["repo_id"],
repo_type=self.spec["repo_type"],
revision=self.spec["revision"],
filename=name,
token=HF_TOKEN,
cache_dir=str(self.cache_dir),
)
def _future_for(self, name):
future = self._futures.get(name)
if future is None:
future = self._download_pool.submit(self._download, name)
self._futures[name] = future
return future
def _open(self, name):
if name in self._mapped:
self._mapped.move_to_end(name)
return self._mapped[name]
future = self._futures.pop(name, None)
path = future.result() if future is not None else self._download(name)
byte_count = Path(path).stat().st_size
if byte_count % self.dtype.itemsize:
raise ValueError(f"Malformed shard byte length: {name}")
token_count = byte_count // self.dtype.itemsize
data = (
None
if token_count < SEQUENCE_LENGTH + 1
else np.memmap(path, mode="r", dtype=self.dtype)
)
self._mapped[name] = (data, token_count)
self._mapped.move_to_end(name)
while len(self._mapped) > MAPPED_SHARD_CACHE:
self._mapped.popitem(last=False)
print(
f"Shard ready: {name} | {token_count:,} actual tokens",
flush=True,
)
return data, token_count
def _schedule_next_for_source(self, source):
c = self.cursor["sources"][source]
epoch = int(c["epoch"])
position = int(c["shard_position"]) + 1
if position == len(self.source_files[source]):
position = 0
epoch += 1
index = int(self._shard_order(source, epoch)[position])
name = self.source_files[source][index]
if name not in self._mapped and name not in self._futures:
self._future_for(name)
def _advance_shard(self, source):
c = self.cursor["sources"][source]
c["sequence_position"] = 0
c["shard_position"] += 1
if c["shard_position"] == len(self.source_files[source]):
c["shard_position"] = 0
c["epoch"] += 1
self._active.pop(source, None)
def _activate(self, source):
# At most one source-epoch worth of shards looking for a usable shard.
for _ in range(len(self.source_files[source]) + 1):
c = self.cursor["sources"][source]
epoch = int(c["epoch"])
position = int(c["shard_position"])
key = (source, epoch, position)
active = self._active.get(source)
if active is not None and active["key"] == key:
return active
index = int(self._shard_order(source, epoch)[position])
name = self.source_files[source][index]
data, token_count = self._open(name)
sequence_count = max(0, (token_count - 1) // SEQUENCE_LENGTH)
if sequence_count == 0:
if c["sequence_position"] != 0:
raise ValueError("Saved cursor points into a tiny shard.")
self._advance_shard(source)
continue
if c["sequence_position"] > sequence_count:
raise ValueError(
f"Saved cursor exceeds actual shard length for source {source!r}."
)
source_id = self.SOURCE_NAMES.index(source)
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, source_id, epoch, index, 202])
)
sequence_order = rng.permutation(sequence_count)
active = {
"key": key,
"name": name,
"data": data,
"sequence_order": sequence_order,
}
self._active[source] = active
self._schedule_next_for_source(source)
return active
raise ValueError(f"No shards for source {source!r} contain a complete sequence.")
def _copy_one(self, source, destination):
while True:
active = self._activate(source)
c = self.cursor["sources"][source]
position = int(c["sequence_position"])
order = active["sequence_order"]
if position < len(order):
break
self._advance_shard(source)
sample = int(order[position])
start = sample * SEQUENCE_LENGTH
values = active["data"][start:start + SEQUENCE_LENGTH + 1]
if len(values) != SEQUENCE_LENGTH + 1:
raise RuntimeError("Internal shard bounds invariant failed.")
np.copyto(destination, values, casting="safe")
c["sequence_position"] += 1
def batch(self):
storage = torch.empty(
(GLOBAL_MICRO_BATCH_SIZE, SEQUENCE_LENGTH + 1),
dtype=torch.long,
pin_memory=True,
)
array = storage.numpy()
counts = {name: 0 for name in self.SOURCE_NAMES}
for row in range(GLOBAL_MICRO_BATCH_SIZE):
global_sequence = int(self.cursor["global_sequence"])
source = self._source_for_global(global_sequence)
self._copy_one(source, array[row])
self.cursor["global_sequence"] = global_sequence + 1
counts[source] += 1
maximum = int(array.max())
if maximum >= int(self.spec["vocab_size"]):
raise ValueError(
f"Token ID {maximum} is outside vocabulary {self.spec['vocab_size']}."
)
return storage, self.state_dict()
def close(self):
self._download_pool.shutdown(wait=False, cancel_futures=True)
self._active.clear()
self._mapped.clear()
self._futures.clear()
class DataPipelineError(RuntimeError):
pass
class PrefetchedStream:
def __init__(self, source):
self.source = source
self.queue = queue.Queue(maxsize=max(1, PREFETCH_BATCHES))
self.stopped = threading.Event()
self.wait_seconds = 0.0
self.thread = threading.Thread(
target=self._worker,
name="token-prefetch",
daemon=True,
)
self.thread.start()
def _put(self, item):
while not self.stopped.is_set():
try:
self.queue.put(item, timeout=0.2)
return True
except queue.Full:
pass
return False
def _worker(self):
try:
while not self.stopped.is_set():
storage, cursor = self.source.batch()
if not self._put((True, storage, cursor)):
return
except Exception as exc:
self._put((False, exc, None))
finally:
self.source.close()
def batch(self):
started = time.monotonic()
while True:
try:
ok, value, cursor = self.queue.get(timeout=0.5)
break
except queue.Empty:
if not self.thread.is_alive():
raise DataPipelineError("Data producer exited.")
self.wait_seconds += time.monotonic() - started
if not ok:
raise DataPipelineError("Data producer failed") from value
# One contiguous H2D copy rather than two overlapping transfers.
storage = value.to("cuda", non_blocking=True)
return storage[:, :-1], storage[:, 1:], cursor
def close(self):
self.stopped.set()
self.thread.join(timeout=2.0)
def distributed_batch(stream):
"""Fetch one global batch on rank zero and give every rank a disjoint slice."""
local_batch = GLOBAL_MICRO_BATCH_SIZE // WORLD_SIZE
failure = None
status = None
if RANK == 0:
try:
x, y, cursor = stream.batch()
packed = torch.stack((x, y), dim=0)
status = "ok"
except Exception as exc:
failure = exc
status = "oom" if isinstance(exc, torch.cuda.OutOfMemoryError) else "failed"
# Never leave peers waiting for a batch when its producer has failed. All
# ranks take the same recovery path, using the last committed data cursor.
status = broadcast_object(status)
if status == "oom":
raise torch.cuda.OutOfMemoryError("Rank-zero batch preparation ran out of memory.") from failure
if status != "ok":
raise DataPipelineError("Rank-zero batch preparation failed.") from failure
if RANK != 0:
packed = torch.empty(
(2, GLOBAL_MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device=DEVICE,
dtype=torch.long,
)
cursor = None
if dist.is_initialized():
dist.broadcast(packed, src=0)
start = RANK * local_batch
end = start + local_batch
return packed[0, start:end], packed[1, start:end], cursor
# ===========================================================================
# Checkpoint storage
# ===========================================================================
def optimizer_layout(optimizer, model):
names = {id(p): name for name, p in model.named_parameters()}
layout = []
for group in optimizer.param_groups:
saved = {k: v for k, v in group.items() if k != "params"}
saved["param_names"] = [names[id(p)] for p in group["params"]]
layout.append(saved)
return layout
def training_settings():
return {
"peak_lr": PEAK_LR,
"min_lr": MIN_LR,
"warmup_steps": WARMUP_STEPS,
"max_steps": MAX_STEPS,
"weight_decay": WEIGHT_DECAY,
"grad_clip": GRAD_CLIP,
"router_aux_coef": ROUTER_AUX_COEF,
"router_z_coef": ROUTER_Z_COEF,
"router_specialization_coef": ROUTER_SPECIALIZATION_COEF,
"spec_warmup_steps": SPEC_WARMUP_STEPS,
"train_epochs": TRAIN_EPOCHS,
"data_target_tokens": DATA_TARGET_TOKENS,
"requested_train_tokens": REQUESTED_TRAIN_TOKENS,
"train_tokens": TRAIN_TOKENS,
"global_micro_batch_size": GLOBAL_MICRO_BATCH_SIZE,
"micro_batch_size": MICRO_BATCH_SIZE,
"world_size": WORLD_SIZE,
"grad_accum_steps": GRAD_ACCUM_STEPS,
"tokens_per_update": TOKENS_PER_UPDATE,
"data_stream_version": DATA_STREAM_VERSION,
"run_id": RUN_ID,
}
def checkpoint_valid(path):
try:
path = Path(path)
manifest = json.loads((path / "manifest.json").read_text())
if not isinstance(manifest, dict):
return False
if (
manifest.get("run_id") != RUN_ID
or manifest.get("format_version") != 4
or manifest.get("checkpoint_format") != "ddp_full_state"
):
return False
shards = []
for key in ("model_shards", "optimizer_shards"):
names = manifest.get(key)
if not isinstance(names, list) or not names:
return False
if any(not isinstance(name, str) for name in names):
return False
shards.extend(safe_relative_path(name) for name in names)
if len(shards) != len(set(shards)):
return False
files = (
shards
+ ["training_state.pt", "tokenizer/tokenizer_config.json"]
)
return all(
(path / safe_relative_path(name)).is_file()
for name in files
)
except (OSError, ValueError, KeyError, TypeError):
return False
def latest_local_checkpoint():
if RESUME_LOCAL_PATH is not None:
path = Path(RESUME_LOCAL_PATH)
if not checkpoint_valid(path):
raise ValueError(f"Incomplete local checkpoint: {path}")
return path
root = WORK / "checkpoints"
candidates = []
if root.exists():
for path in root.glob("step-*"):
if not path.is_dir() or not checkpoint_valid(path):
continue
manifest = json.loads((path / "manifest.json").read_text())
progress = manifest.get("progress")
# Original format-1 saves lack progress in their manifest.
if progress is None:
state = torch.load(
path / "training_state.pt",
map_location="cpu",
weights_only=False,
)
progress = state["progress"]
del state
candidates.append((
int(progress["step"]),
int(progress.get("tokens", 0)),
path.stat().st_mtime_ns,
path,
))
return max(candidates, default=(None, None, None, None))[-1]
def checkpoint_progress(path):
manifest = json.loads((path / "manifest.json").read_text())
if "progress" in manifest:
return manifest["progress"]
state = torch.load(
path / "training_state.pt",
map_location="cpu",
weights_only=False,
)
return state["progress"]
def capture_local_checkpoint(model, optimizer, tokenizer, progress, data_state):
"""Capture an owned CPU snapshot on rank zero at an optimizer boundary."""
module = model.module if hasattr(model, "module") else model
torch.cuda.synchronize(DEVICE)
snapshot = {
"full_model": cpu_tree(module.state_dict()),
"full_optimizer": cpu_tree(optimizer.state_dict()),
"tokenizer": copy.deepcopy(tokenizer),
"progress": copy.deepcopy(progress),
"data_state": copy.deepcopy(cpu_tree(data_state)),
"cfg": copy.deepcopy(cpu_tree(module.cfg)),
"optimizer_groups": copy.deepcopy(cpu_tree(optimizer_layout(optimizer, module))),
}
snapshot["training_state"] = {
"progress": snapshot["progress"],
"optimizer_groups": snapshot["optimizer_groups"],
"data_state": snapshot["data_state"],
"training_settings": copy.deepcopy(training_settings()),
"versions": {
"torch": str(torch.__version__),
"optimizer": "torch.optim.AdamW",
},
**copy.deepcopy(cpu_tree(capture_rng())),
}
return snapshot
class AsyncCheckpointWriter:
"""Single-flight CPU-only serializer; completed work must be collected."""
def __init__(self):
self.future = None
self.executor = ThreadPoolExecutor(
max_workers=1, thread_name_prefix="checkpoint-write"
)
def submit(self, snapshot):
if self.future is not None:
raise RuntimeError("Previous local checkpoint must be collected first.")
self.future = self.executor.submit(_write_local_checkpoint, **snapshot)
def busy(self):
return self.future is not None and not self.future.done()
def failure(self):
if self.future is not None and self.future.done():
return self.future.exception()
return None
def result(self):
if self.future is None or self.busy():
return None
return self.wait()
def wait(self):
if self.future is None:
return None
future = self.future
try:
return future.result()
finally:
self.future = None
def close(self):
try:
self.wait()
finally:
self.executor.shutdown(wait=True)
def save_local_checkpoint(
model, optimizer, tokenizer, progress, data_state
):
final = None
failure = None
error = None
if RANK == 0:
try:
snapshot = capture_local_checkpoint(
model, optimizer, tokenizer, progress, data_state
)
final = _write_local_checkpoint(**snapshot)
except BaseException as exc:
failure = exc
error = f"{type(exc).__name__}: {exc}"
# Collection, CPU copying, and filesystem errors all reach the same single
# status broadcast. No state-dict collective or success-only barrier.
error = broadcast_object(error)
if error is not None:
if failure is not None:
raise failure
raise RuntimeError(f"Rank-zero checkpoint save failed: {error}") from failure
return final
def _write_local_checkpoint(
full_model, full_optimizer, tokenizer, progress, data_state, cfg,
optimizer_groups, training_state,
):
root = WORK / "checkpoints"
root.mkdir(parents=True, exist_ok=True)
name = f"step-{progress['step']:09d}-{uuid.uuid4().hex[:8]}"
temporary = root / (".building-" + name)
final = root / name
required = nested_bytes(full_model) + nested_bytes(full_optimizer) + 2 * 1024**3
free = shutil.disk_usage(root).free
if free < required:
raise RuntimeError(
f"Not enough free disk for another checkpoint. "
f"Need about {required / 1024**3:.1f} GiB; "
f"have {free / 1024**3:.1f} GiB."
)
temporary.mkdir(exist_ok=False)
manifest = {
"format_version": 4,
"checkpoint_format": "ddp_full_state",
"run_id": RUN_ID,
"architecture": cfg,
"progress": copy.deepcopy(progress),
"model_shards": [],
"optimizer_shards": [],
}
try:
for i, shard in enumerate(shard_items(full_model.items())):
filename = f"model-{i:04d}.safetensors"
save_file(shard, str(temporary / filename))
manifest["model_shards"].append(filename)
# Splitting the top-level optimizer dictionary would put every moment
# tensor in one enormous "state" shard. Split per parameter instead;
# only the first file owns param_groups, including empty optimizer state.
optimizer_shards = shard_items(sorted(full_optimizer["state"].items()))
for i, shard in enumerate(optimizer_shards):
filename = f"optimizer-{i:04d}.pt"
payload = {"state": shard}
if i == 0:
payload["param_groups"] = full_optimizer["param_groups"]
torch.save(payload, temporary / filename)
manifest["optimizer_shards"].append(filename)
if not manifest["optimizer_shards"]:
filename = "optimizer-0000.pt"
torch.save(full_optimizer, temporary / filename)
manifest["optimizer_shards"].append(filename)
torch.save(training_state, temporary / "training_state.pt")
tokenizer.save_pretrained(temporary / "tokenizer")
(temporary / "config.json").write_text(
json.dumps(cfg, indent=2)
)
(temporary / "README.md").write_text(
f"# {cfg.get('model_name', 'Custom PyTorch model')} training checkpoint\n\n"
"Custom PyTorch architecture, not Transformers AutoModel.\n"
"Format 4: portable DDP full model/optimizer state.\n"
"Only load trusted optimizer/training pickle files.\n"
)
# Completion marker written last.
atomic_json(temporary / "manifest.json", manifest)
for path in temporary.rglob("*"):
if path.is_file():
with path.open("rb") as f:
os.fsync(f.fileno())
os.replace(temporary, final)
atomic_json(WORK / "latest-local.json", {
"path": str(final.resolve()),
"step": progress["step"],
"tokens": progress["tokens"],
})
return final
except BaseException:
shutil.rmtree(temporary, ignore_errors=True)
raise
def prune_local_checkpoints(keep):
keep = Path(keep).resolve()
root = WORK / "checkpoints"
if root.exists():
for path in root.glob("step-*"):
if path.is_dir() and path.resolve() != keep:
shutil.rmtree(path)
def _load_checkpoint_shards(path, manifest, key, load_shard):
filenames = manifest.get(key)
if not isinstance(filenames, list) or not filenames:
raise ValueError(f"Checkpoint manifest requires a non-empty {key} list.")
is_optimizer = key == "optimizer_shards"
merged = {"state": {}} if is_optimizer else {}
has_optimizer_state = False
seen = set()
for filename in filenames:
if not isinstance(filename, str):
raise ValueError(f"Invalid filename in checkpoint {key}: {filename!r}")
filename = safe_relative_path(filename)
if filename in seen:
raise ValueError(f"Duplicate checkpoint shard in {key}: {filename}")
seen.add(filename)
shard = load_shard(path / filename)
if not isinstance(shard, dict):
raise ValueError(f"Checkpoint shard {filename} must contain a dictionary.")
if is_optimizer:
if not shard or shard.keys() - {"state", "param_groups"}:
raise ValueError(f"Invalid optimizer shard contents: {filename}")
if "state" in shard:
has_optimizer_state = True
states = shard["state"]
if not isinstance(states, dict):
raise ValueError(f"Optimizer state in {filename} must be a dictionary.")
duplicates = merged["state"].keys() & states.keys()
if duplicates:
raise ValueError(
f"Duplicate optimizer state entries in {filename}: "
f"{sorted(map(str, duplicates))}"
)
merged["state"].update(states)
if "param_groups" in shard:
if "param_groups" in merged:
raise ValueError("Checkpoint has multiple optimizer param_groups entries.")
if not isinstance(shard["param_groups"], list):
raise ValueError("Optimizer param_groups must be a list.")
merged["param_groups"] = shard["param_groups"]
continue
duplicates = merged.keys() & shard.keys()
if duplicates:
raise ValueError(
f"Duplicate keys in checkpoint {key} shard {filename}: "
f"{sorted(map(str, duplicates))}"
)
merged.update(shard)
if is_optimizer and "param_groups" not in merged:
raise ValueError("Checkpoint optimizer shards lack param_groups.")
if is_optimizer and not has_optimizer_state:
raise ValueError("Checkpoint optimizer shards lack state.")
return merged
def restore_checkpoint(path, model, optimizer):
def agree_on_failure(failure, phase):
local_error = (
f"rank {RANK}: {type(failure).__name__}: {failure}"
if failure is not None else None
)
errors = [local_error]
if dist.is_initialized():
errors = [None] * dist.get_world_size()
dist.all_gather_object(errors, local_error)
error = broadcast_object(
"; ".join(item for item in errors if item) or None
)
if error is not None:
if failure is not None:
raise failure
raise RuntimeError(f"Checkpoint {phase} failed: {error}")
failure = None
model_state = optimizer_state = state = None
try:
path = Path(path)
module = model.module if hasattr(model, "module") else model
manifest = json.loads((path / "manifest.json").read_text())
if not isinstance(manifest, dict):
raise ValueError("Checkpoint manifest must be a dictionary.")
if (
manifest.get("format_version") != 4
or manifest.get("checkpoint_format") != "ddp_full_state"
):
raise ValueError("This DDP trainer only accepts format-4 ddp_full_state checkpoints.")
if manifest.get("run_id") != RUN_ID:
raise ValueError("Checkpoint belongs to an older/incompatible model run.")
if manifest["architecture"] != module.cfg:
raise ValueError("Checkpoint architecture does not match.")
model_state = _load_checkpoint_shards(
path, manifest, "model_shards",
lambda shard: load_file(str(shard), device="cpu"),
)
optimizer_state = _load_checkpoint_shards(
path, manifest, "optimizer_shards",
lambda shard: torch.load(shard, map_location="cpu", weights_only=False),
)
state = torch.load(
path / "training_state.pt", map_location="cpu", weights_only=False,
)
expected_state = module.state_dict()
if model_state.keys() != expected_state.keys():
raise ValueError("Checkpoint model parameter/buffer names do not match.")
for name, tensor in model_state.items():
if not torch.is_tensor(tensor) or tensor.shape != expected_state[name].shape:
raise ValueError(f"Checkpoint model tensor shape does not match: {name}")
del expected_state
# Optimizer IDs are positional, not model names. Validate the complete
# group/name layout before PyTorch maps saved IDs to live parameters.
# Unused experts may legitimately have no moments yet; require a subset
# of registered IDs, not an optimizer state for every parameter.
saved_layout = state.get("optimizer_groups")
current_layout = optimizer_layout(optimizer, module)
saved_groups = optimizer_state["param_groups"]
if (
not isinstance(saved_layout, list)
or len(saved_layout) != len(current_layout)
or len(saved_groups) != len(current_layout)
):
raise ValueError("Checkpoint optimizer parameter-group layout does not match.")
parameter_ids = set()
for saved, current, group in zip(saved_layout, current_layout, saved_groups):
if not isinstance(saved, dict) or not isinstance(group, dict):
raise ValueError("Invalid checkpoint optimizer parameter group.")
names = saved.get("param_names")
ids = group.get("params")
if names != current["param_names"] or not isinstance(ids, list):
raise ValueError("Checkpoint optimizer parameter names/order do not match.")
if len(ids) != len(names):
raise ValueError("Checkpoint optimizer positional layout does not match.")
expected_ids = list(range(len(parameter_ids), len(parameter_ids) + len(names)))
if ids != expected_ids:
raise ValueError("Checkpoint optimizer positional parameter IDs do not match.")
if "param_names" in group and group["param_names"] != names:
raise ValueError("Conflicting optimizer parameter-name metadata.")
for parameter_id in ids:
if type(parameter_id) is not int or parameter_id < 0 or parameter_id in parameter_ids:
raise ValueError("Invalid or duplicate checkpoint optimizer parameter ID.")
parameter_ids.add(parameter_id)
for parameter_id, entry in optimizer_state["state"].items():
if type(parameter_id) is not int or parameter_id not in parameter_ids:
raise ValueError("Checkpoint optimizer state contains an unknown parameter ID.")
if not isinstance(entry, dict):
raise ValueError("Checkpoint optimizer parameter state must be a dictionary.")
previous = state.get("training_settings", {})
current_settings = training_settings()
for key in (
"peak_lr", "min_lr", "warmup_steps", "max_steps",
"router_aux_coef", "router_z_coef", "router_specialization_coef",
"spec_warmup_steps", "train_epochs", "data_target_tokens",
"requested_train_tokens", "train_tokens", "tokens_per_update",
"data_stream_version", "run_id",
):
if key in previous and previous[key] != current_settings[key]:
raise ValueError(
f"Resume setting changed: {key}: "
f"{previous[key]} -> {current_settings[key]}. "
"Keep the original schedule/objective for this resume."
)
except BaseException as exc:
failure = exc
# Disk reads and validation are side-effect-free. Nobody mutates live state
# or starts training until every rank has successfully completed this phase.
agree_on_failure(failure, "read/validation")
failure = None
try:
module.load_state_dict(model_state)
del model_state
optimizer.load_state_dict(optimizer_state)
del optimizer_state
restore_rng(state)
except BaseException as exc:
failure = exc
# Also coordinate device/allocation and optimizer-loader failures so peers
# cannot enter the next DDP forward while one rank is unwinding its restore.
agree_on_failure(failure, "application")
rank0_print("Restored versions:", state.get("versions"), flush=True)
return state
# ===========================================================================
# Hugging Face upload manager
# ===========================================================================
class HFCheckpoints:
def __init__(self):
self.api = HfApi(token=HF_TOKEN)
self.api.create_repo(
repo_id=HF_REPO_ID,
repo_type="model",
private=HF_PRIVATE,
exist_ok=True,
)
if USE_LARGE_FOLDER_UPLOAD and not hasattr(
self.api, "upload_large_folder"
):
raise RuntimeError(
"Upgrade huggingface_hub, or set "
"USE_LARGE_FOLDER_UPLOAD = False."
)
self.previous_remote = None
self.pointer = None
self.revision = None
self.future = None
self.last_upload_seconds = None
self.executor = ThreadPoolExecutor(
max_workers=1, thread_name_prefix="checkpoint-upload"
)
def _remote_checkpoint_inventory(self):
"""
Return current repo revision + file inventory.
We intentionally inspect the actual repo tree instead of trusting
latest.json. latest.json is only a convenience pointer and may be stale
after a manually deleted checkpoint or an interrupted upload.
"""
info = self._api_retry(
"repo info",
lambda: self.api.repo_info(
repo_id=HF_REPO_ID,
repo_type="model",
),
)
self.revision = info.sha
files = self._api_retry(
"remote checkpoint inventory",
lambda: self.api.list_repo_files(
repo_id=HF_REPO_ID,
repo_type="model",
revision=self.revision,
),
)
return self.revision, set(files)
def _remote_manifest_pointer(self, folder, files):
"""
Validate one remote checkpoint using its manifest and the repo file
inventory. Returns a pointer dict when COMPLETE, else None.
"""
folder = safe_relative_path(folder).rstrip("/")
manifest_name = f"{folder}/manifest.json"
if manifest_name not in files:
return None
try:
manifest_path = hf_hub_download(
repo_id=HF_REPO_ID,
repo_type="model",
revision=self.revision,
filename=manifest_name,
token=HF_TOKEN,
cache_dir=str(WORK / "hf_metadata"),
)
manifest = json.loads(Path(manifest_path).read_text())
if (
manifest.get("run_id") != RUN_ID
or manifest.get("format_version") != 4
or manifest.get("checkpoint_format") != "ddp_full_state"
):
return None
shards = []
for key in ("model_shards", "optimizer_shards"):
names = manifest.get(key)
if not isinstance(names, list) or not names:
return None
if any(not isinstance(name, str) for name in names):
return None
shards.extend(safe_relative_path(name) for name in names)
if len(shards) != len(set(shards)):
return None
required = (
shards
+ [
"training_state.pt",
"config.json",
"tokenizer/tokenizer_config.json",
"manifest.json",
]
)
required_remote = {
f"{folder}/{safe_relative_path(name)}"
for name in required
}
if not required_remote.issubset(files):
return None
progress = manifest.get("progress") or {}
match = re.match(r"^checkpoints/step-(\d+)-[^/]+$", folder)
parsed_step = int(match.group(1)) if match else 0
return {
"checkpoint": folder,
"step": int(progress.get("step", parsed_step)),
"tokens": int(progress.get("tokens", 0)),
}
except Exception as exc:
print(
f"Warning: ignoring invalid remote checkpoint {folder}: {exc}",
flush=True,
)
return None
def find_latest_complete_remote(self):
"""
Scan checkpoints/* and return the newest COMPLETE checkpoint.
This is the source of truth for resume. It recovers automatically when
latest.json points at a deleted, partial, or older checkpoint.
"""
_, files = self._remote_checkpoint_inventory()
folders = set()
pattern = re.compile(r"^(checkpoints/step-(\d+)-[^/]+)/manifest\.json$")
for name in files:
match = pattern.match(name)
if match:
folders.add((int(match.group(2)), match.group(1)))
# Check highest step first; tokens break ties after reading manifest.
candidates = []
for _, folder in sorted(folders, reverse=True):
pointer = self._remote_manifest_pointer(folder, files)
if pointer is not None:
candidates.append(pointer)
if not candidates:
return None
return max(
candidates,
key=lambda p: (int(p["step"]), int(p.get("tokens", 0))),
)
def _publish_pointer(self, pointer, message):
payload = dict(pointer)
payload["updated_utc"] = datetime.datetime.now(
datetime.timezone.utc
).isoformat()
self._api_retry(
"latest pointer repair/publish",
lambda: self.api.upload_file(
repo_id=HF_REPO_ID,
repo_type="model",
path_or_fileobj=json.dumps(payload, indent=2).encode(),
path_in_repo="latest.json",
commit_message=message,
),
)
self.pointer = payload
self.previous_remote = payload["checkpoint"]
return payload
def resolve_resume_pointer(self, repair=True):
"""
Resolve a trustworthy remote resume target.
Rules:
* Never trust latest.json blindly.
* Pick the highest COMPLETE checkpoint actually present on the Hub.
* If latest.json is stale/deleted/older, repair it automatically.
"""
raw_pointer = None
try:
raw_pointer = self.read_latest()
except Exception as exc:
print("Warning: latest.json could not be read:", exc, flush=True)
complete = self.find_latest_complete_remote()
if complete is None:
self.pointer = None
self.previous_remote = None
return None
raw_key = (
int(raw_pointer.get("step", -1)),
int(raw_pointer.get("tokens", -1)),
raw_pointer.get("checkpoint"),
) if raw_pointer else (-1, -1, None)
complete_key = (
int(complete["step"]),
int(complete.get("tokens", 0)),
complete["checkpoint"],
)
# latest.json can point to a deleted checkpoint at the same numeric
# step, so checkpoint path is part of the comparison.
if raw_key != complete_key:
print(
"latest.json is stale or not the newest complete checkpoint.",
flush=True,
)
print(
"Recovered remote checkpoint:",
complete["checkpoint"],
f"(step={complete['step']}, tokens={complete.get('tokens', 0)})",
flush=True,
)
if repair:
complete = self._publish_pointer(
complete,
message=f"Repair latest pointer to step {complete['step']}",
)
print("Repaired latest.json automatically.", flush=True)
else:
self.pointer = complete
self.previous_remote = complete["checkpoint"]
else:
self.pointer = complete
self.previous_remote = complete["checkpoint"]
return self.pointer
def read_latest(self):
info = self.api.repo_info(
repo_id=HF_REPO_ID, repo_type="model"
)
self.revision = info.sha
try:
path = hf_hub_download(
repo_id=HF_REPO_ID,
repo_type="model",
revision=self.revision,
filename="latest.json",
token=HF_TOKEN,
cache_dir=str(WORK / "hf_metadata"),
)
except EntryNotFoundError:
return None
pointer = json.loads(Path(path).read_text())
remote = safe_relative_path(pointer["checkpoint"])
if not remote.startswith("checkpoints/"):
raise ValueError("Unexpected checkpoint path in latest.json.")
self.pointer = pointer
self.previous_remote = remote
return pointer
def download_latest(self):
if self.pointer is None:
return None
for attempt in range(2):
remote = self.pointer["checkpoint"]
print("Downloading checkpoint:", remote, flush=True)
root = snapshot_download(
repo_id=HF_REPO_ID,
repo_type="model",
revision=self.revision,
token=HF_TOKEN,
allow_patterns=[f"{remote}/*"],
local_dir=str(WORK / "download"),
max_workers=8,
)
path = Path(root) / remote
if checkpoint_valid(path):
return path
if attempt == 0:
print(
"Remote checkpoint changed/disappeared during download; "
"rescanning the Hub once.",
flush=True,
)
pointer = self.resolve_resume_pointer(repair=True)
if pointer is None:
break
raise RuntimeError(
"No complete remote checkpoint could be downloaded after rescan."
)
def busy(self):
return self.future is not None and not self.future.done()
def wait(self):
if self.future is None:
return True
future = self.future
self.future = None
try:
future.result()
return True
except Exception as exc:
print(
f"WARNING: upload failed: {exc}\n"
"The full local checkpoint is still available.",
flush=True,
)
return False
def _api_retry(self, label, fn):
"""Retry Hub API calls on rate limits and transient server failures."""
delay = HF_API_RETRY_BASE_SECONDS
for attempt in range(1, HF_API_MAX_RETRIES + 1):
try:
return fn()
except Exception as exc:
status = getattr(getattr(exc, "response", None), "status_code", None)
transient = status == 429 or (status is not None and 500 <= status < 600)
if not transient or attempt >= HF_API_MAX_RETRIES:
raise
retry_after = None
response = getattr(exc, "response", None)
if response is not None:
try:
retry_after = float(response.headers.get("Retry-After", ""))
except (TypeError, ValueError):
retry_after = None
sleep_for = retry_after if retry_after is not None else delay
sleep_for = min(max(1.0, sleep_for), HF_API_RETRY_MAX_SECONDS)
print(
f"HF {label}: HTTP {status}; retry {attempt}/"
f"{HF_API_MAX_RETRIES} in {sleep_for:.1f}s",
flush=True,
)
time.sleep(sleep_for)
delay = min(delay * 2.0, HF_API_RETRY_MAX_SECONDS)
def estimate_upload_seconds(self, size):
estimate = size / (ESTIMATED_UPLOAD_MB_PER_SECOND * 1_000_000)
if self.last_upload_seconds is not None:
estimate = max(estimate, self.last_upload_seconds * 1.25)
return estimate
def upload_async(self, path, progress):
if self.future is not None:
raise RuntimeError("Previous upload must be joined first.")
self.future = self.executor.submit(
self._upload, Path(path), copy.deepcopy(progress)
)
def _upload(self, path, progress):
started = time.monotonic()
remote = f"checkpoints/{path.name}"
# Never silently move a mature run backwards. This guard runs even if
# somebody accidentally leaves START_FROM_SCRATCH=True.
if not ALLOW_REMOTE_POINTER_ROLLBACK:
try:
current = self.find_latest_complete_remote()
except Exception as exc:
current = None
print(
"WARNING: could not verify remote monotonicity before "
f"upload: {exc}",
flush=True,
)
if current is not None:
new_key = (
int(progress["step"]),
int(progress.get("tokens", 0)),
)
current_key = (
int(current["step"]),
int(current.get("tokens", 0)),
)
if new_key < current_key:
print(
"REMOTE ROLLBACK GUARD: refusing to publish/upload "
f"step {new_key[0]} because the Hub already has a "
f"newer complete checkpoint at step {current_key[0]} "
f"({current['checkpoint']}). Local checkpoint kept.",
flush=True,
)
self.pointer = current
self.previous_remote = current["checkpoint"]
return
previous = self.previous_remote
staging = WORK / "upload-staging" / path.name
# Keep failed large-upload staging so a restart can reuse its
# upload metadata. Before a different checkpoint is uploaded,
# remove obsolete staging/hardlinks to bound local disk usage.
staging_root = WORK / "upload-staging"
staging_root.mkdir(parents=True, exist_ok=True)
for old in staging_root.iterdir():
if old.is_dir() and old != staging:
shutil.rmtree(old)
if USE_LARGE_FOLDER_UPLOAD:
destination = staging / remote
if not destination.exists():
destination.parent.mkdir(parents=True, exist_ok=True)
try:
shutil.copytree(path, destination, copy_function=os.link)
except OSError as exc:
shutil.rmtree(staging, ignore_errors=True)
raise RuntimeError(
"Hardlink upload staging failed. Keep WORK_DIR on "
"one filesystem or set USE_LARGE_FOLDER_UPLOAD=False."
) from exc
# upload_large_folder has no path_in_repo argument. The
# hardlinked staging tree supplies the repository hierarchy.
self._api_retry(
"large checkpoint upload",
lambda: self.api.upload_large_folder(
repo_id=HF_REPO_ID,
repo_type="model",
folder_path=str(staging),
num_workers=UPLOAD_WORKERS,
),
)
else:
self._api_retry(
"checkpoint upload",
lambda: self.api.upload_folder(
repo_id=HF_REPO_ID,
repo_type="model",
folder_path=str(path),
path_in_repo=remote,
commit_message=f"{MODEL_NAME}: step {progress['step']}",
),
)
# Publish only after the entire checkpoint upload completes.
pointer = {
"checkpoint": remote,
"step": progress["step"],
"tokens": progress["tokens"],
"updated_utc": datetime.datetime.now(
datetime.timezone.utc
).isoformat(),
}
self._api_retry(
"latest pointer publish",
lambda: self.api.upload_file(
repo_id=HF_REPO_ID,
repo_type="model",
path_or_fileobj=json.dumps(pointer, indent=2).encode(),
path_in_repo="latest.json",
commit_message=f"Publish step {progress['step']}",
),
)
self.previous_remote = remote
self.pointer = pointer
print("HF checkpoint published:", remote, flush=True)
cleaned_previous = False
if (
KEEP_ONLY_LATEST_REMOTE_FOLDER
and previous
and previous != remote
):
try:
self._api_retry(
"old checkpoint deletion",
lambda: self.api.delete_folder(
repo_id=HF_REPO_ID,
repo_type="model",
path_in_repo=previous,
commit_message="Remove previous checkpoint folder",
),
)
cleaned_previous = True
print("Removed old remote checkpoint:", previous, flush=True)
except Exception as exc:
# Do NOT squash if deletion failed: that could preserve the old
# checkpoint in the new root commit and defeat the cleanup.
print("WARNING: old remote folder cleanup failed:", exc, flush=True)
if SUPER_SQUASH_AFTER_REMOTE_CLEANUP and cleaned_previous:
try:
self._api_retry(
"history squash",
lambda: self.api.super_squash_history(
repo_id=HF_REPO_ID,
repo_type="model",
branch="main",
commit_message=(
f"Rolling checkpoint storage: keep step "
f"{progress['step']} only"
),
),
)
print(
"HF history super-squashed; old checkpoint LFS history "
"is no longer retained by main.",
flush=True,
)
except Exception as exc:
print(
"WARNING: history squash failed. The old folder is gone "
"from main, but historical LFS blobs may still count "
"toward storage:",
exc,
flush=True,
)
shutil.rmtree(staging, ignore_errors=True)
self.last_upload_seconds = time.monotonic() - started
print(
f"Upload finished in {self.last_upload_seconds / 60:.1f} min. "
"Local checkpoint retained.",
flush=True,
)
def close(self):
self.wait()
self.executor.shutdown(wait=True)
def choose_resume(hub):
"""Newest complete checkpoint wins; fresh start only when emptiness is verified."""
explicit = None
if RESUME_LOCAL_PATH is not None:
candidate = Path(RESUME_LOCAL_PATH)
if checkpoint_valid(candidate):
explicit = candidate
print("Using explicit complete local checkpoint:", candidate, flush=True)
else:
print(
"WARNING: RESUME_LOCAL_PATH is missing/incomplete:", candidate,
"— falling back to normal local + remote discovery.",
flush=True,
)
if explicit is not None:
try:
hub.resolve_resume_pointer(repair=True)
except Exception as exc:
print("Warning: remote lookup/repair failed:", exc, flush=True)
return explicit, False
local = latest_local_checkpoint()
try:
pointer = hub.resolve_resume_pointer(repair=True)
remote_lookup_ok = True
except Exception as exc:
remote_lookup_ok = False
pointer = None
if local is None:
raise RuntimeError(
"No local checkpoint and remote checkpoint discovery failed. "
"Refusing to guess that this is a fresh run."
) from exc
print("Remote lookup failed; using local checkpoint:", exc, flush=True)
if START_FROM_SCRATCH:
if pointer is not None and not ALLOW_FRESH_START_WITH_EXISTING_REMOTE:
raise RuntimeError(
"START_FROM_SCRATCH=True but a complete remote checkpoint exists "
f"at step {pointer['step']} ({pointer['checkpoint']})."
)
print("Intentional fresh start requested.", flush=True)
return None, False
local_key = (-1, -1)
if local is not None:
p = checkpoint_progress(local)
local_key = (int(p["step"]), int(p.get("tokens", 0)))
remote_key = (-1, -1)
if pointer is not None:
remote_key = (int(pointer["step"]), int(pointer.get("tokens", 0)))
if local is not None and local_key >= remote_key:
print(
"Using newest complete local checkpoint:", local,
f"(step={local_key[0]}, tokens={local_key[1]})", flush=True,
)
return local, False
if pointer is not None:
print(
"Using newest complete remote checkpoint:", pointer["checkpoint"],
f"(step={remote_key[0]}, tokens={remote_key[1]})", flush=True,
)
return hub.download_latest(), True
if remote_lookup_ok and local is None:
print("No complete 2B checkpoint exists: starting fresh at step 0.", flush=True)
return None, False
raise RuntimeError("Could not resolve a safe resume/fresh-start state.")
# ===========================================================================
# Training
# ===========================================================================
def request_stop(signum, frame):
global STOP_REQUESTED
STOP_REQUESTED = True
print(
"\nStop requested — finishing this optimizer update, then saving.",
flush=True,
)
def learning_rate(step):
if step <= WARMUP_STEPS:
return PEAK_LR * step / max(1, WARMUP_STEPS)
fraction = min(
1.0,
(step - WARMUP_STEPS) / max(1, MAX_STEPS - WARMUP_STEPS),
)
return MIN_LR + 0.5 * (PEAK_LR - MIN_LR) * (
1.0 + math.cos(math.pi * fraction)
)
def _fresh_start_banner():
print(f"{MODEL_NAME} — DDP BASE PRETRAINING (STRATIFIED MIX v3)", flush=True)
print(f"Checkpoint repo: {HF_REPO_ID}", flush=True)
print(f"Training data: {DATA_REPO_ID}@{DATA_REVISION}/{DATA_CACHE_DIR}", flush=True)
print(
f"Budget: {TRAIN_EPOCHS} nominal epochs over a {DATA_TARGET_TOKENS:,}-token corpus; "
f"requested {REQUESTED_TRAIN_TOKENS:,} tokens.",
flush=True,
)
print(
f"Target: {TRAIN_TOKENS:,} tokens ({TRAIN_UPDATE_COUNT:,} updates); "
f"{REQUESTED_TRAIN_TOKENS - TRAIN_TOKENS:,} final partial-update tokens intentionally unused.",
flush=True,
)
print(
"Nominal token budget, not exact per-example epochs: sources replay with "
"deterministic reshuffling; shard tails are skipped.",
flush=True,
)
def main():
global STOP_REQUESTED, MICRO_BATCH_SIZE, GRAD_ACCUM_STEPS
if WORLD_SIZE != 8:
raise RuntimeError(
f"This run is configured for exactly 8 GPUs; torchrun reported {WORLD_SIZE}."
)
hardware, inventory = validate_hardware()
torch.cuda.set_device(LOCAL_RANK)
if not dist.is_initialized():
dist.init_process_group("nccl", timeout=datetime.timedelta(minutes=30))
if GLOBAL_MICRO_BATCH_SIZE % WORLD_SIZE:
raise ValueError("GLOBAL_MICRO_BATCH_SIZE must be divisible by WORLD_SIZE.")
MICRO_BATCH_SIZE = GLOBAL_MICRO_BATCH_SIZE // WORLD_SIZE
GRAD_ACCUM_STEPS = TOKENS_PER_UPDATE // (
GLOBAL_MICRO_BATCH_SIZE * SEQUENCE_LENGTH
)
if PREFLIGHT_SMOKE:
rank0_print(
"ORION FLAGSHIP 2B DDP — 8-rank preflight smoke "
"(no data stream, no Hugging Face upload)",
flush=True,
)
else:
_fresh_start_banner()
STOP_REQUESTED = False
if not HF_TOKEN and not PREFLIGHT_SMOKE:
raise ValueError(
"Set HF_TOKEN to a Hugging Face token with read access to the data "
"and write access to the checkpoint repo. Do not paste it into this file."
)
if not torch.cuda.is_available():
raise RuntimeError("CUDA GPU required.")
if not torch.cuda.is_bf16_supported():
raise RuntimeError("BF16 GPU support required.")
# Hardware and kernel preflight must complete before checkpoint discovery,
# tokenizer downloads, or construction of a real data stream.
capability = hardware["capability"]
rank0_print(
f"Visible CUDA GPUs: {len(inventory)}; participating ranks: {WORLD_SIZE}\n"
f"{_format_hardware_inventory(inventory)}\n"
f"Active hardware profile: {HARDWARE_PROFILE_NAMES[capability]} "
f"({_format_capability(capability)})",
flush=True,
)
rank0_print("Running bounded CUDA/Triton startup preflight...", flush=True)
if RUN_KERNEL_TESTS:
test_kernels()
select_attention(dict(ARCH))
dist_barrier()
process_start = time.monotonic()
hard_deadline = process_start + SESSION_HOURS * 3600
nominal_train_deadline = process_start + MAX_TRAIN_HOURS * 3600
WORK.mkdir(parents=True, exist_ok=True)
rank0_print(
f"GPU: {hardware['name']} | {hardware['memory_gib']:.1f} GiB\n"
f"SM: {capability[0]}.{capability[1]}\n"
f"PyTorch: {torch.__version__} | CUDA: {torch.version.cuda}\n"
f"Batch: {MICRO_BATCH_SIZE}/GPU × {WORLD_SIZE} GPUs × {GRAD_ACCUM_STEPS} × "
f"{SEQUENCE_LENGTH} = {TOKENS_PER_UPDATE:,} tokens/update",
flush=True,
)
if hardware["memory_gib"] < 85:
rank0_print("Warning: defaults target a 96GB-class GPU.")
random.seed(42)
np.random.seed(42)
torch.manual_seed(42)
torch.backends.cuda.matmul.allow_tf32 = True
torch.set_float32_matmul_precision("high")
torch.backends.cuda.enable_flash_sdp(True)
torch.backends.cuda.enable_mem_efficient_sdp(True)
torch.backends.cuda.enable_math_sdp(False)
hub = None
stream = None
checkpoint_writer = None
old_handlers = {}
try:
if PREFLIGHT_SMOKE:
# Do not instantiate HFCheckpoints: its constructor creates the
# remote repository and the normal path may discover/download/upload
# production checkpoints. The tokenizer read below is still needed
# to construct the production-shape model and is not an upload.
resume_path = None
downloaded = False
resume_payload = {"path": None, "downloaded": False}
else:
hub = HFCheckpoints() if RANK == 0 else None
if RANK == 0:
resume_path, downloaded = choose_resume(hub)
resume_payload = {
"path": str(resume_path) if resume_path is not None else None,
"downloaded": downloaded,
}
else:
resume_payload = None
resume_payload = broadcast_object(resume_payload)
resume_path = Path(resume_payload["path"]) if resume_payload["path"] else None
downloaded = bool(resume_payload["downloaded"])
dist_barrier()
tokenizer_source = (
str(resume_path / "tokenizer")
if resume_path is not None else TOKENIZER_NAME
)
tokenizer = AutoTokenizer.from_pretrained(
tokenizer_source, token=HF_TOKEN or None, use_fast=True
)
if tokenizer.eos_token_id is None:
raise ValueError("Tokenizer requires an EOS token.")
tokenizer.model_max_length = 10**12
cfg = dict(ARCH)
cfg["vocab_size"] = len(tokenizer)
if resume_path is not None:
manifest = json.loads((resume_path / "manifest.json").read_text())
if manifest["architecture"] != cfg:
raise ValueError("Checkpoint architecture/tokenizer mismatch.")
total, active = parameter_counts(cfg)
rank0_print(
f"{MODEL_NAME}: {total / 1e6:.3f}M total parameters; "
f"{active / 1e6:.3f}M approximate active-path compute-equivalent.",
flush=True,
)
if not 1_950_000_000 <= total <= 2_050_000_000:
raise RuntimeError(
f"2B parameter guard failed: expected 1.95–2.05B, got {total:,}."
)
rank0_print("Parameter breakdown:", parameter_count_breakdown(cfg), flush=True)
rank0_print(
f"DDP per-GPU model+grad+Adam FP32 storage: ~{total * 16 / 1024**3:.1f} GiB "
"before activations, autocast cache, communication, and temporaries.",
flush=True,
)
def allocate_model(initialize_weights):
rank0_print("Allocating FP32 T2.1 model...", flush=True)
with torch.device("meta"):
candidate = OrionFlagshipT21(cfg)
candidate.to_empty(device=DEVICE)
for block in candidate.blocks:
block.attn.reset_rope()
if initialize_weights:
candidate.initialize_weights()
candidate.train()
assert sum(p.numel() for p in candidate.parameters()) == total
return DDP(
candidate,
device_ids=[LOCAL_RANK],
output_device=LOCAL_RANK,
find_unused_parameters=True,
broadcast_buffers=False,
gradient_as_bucket_view=True,
)
model = allocate_model(resume_path is None)
# Probe before creating/consuming a real data stream. This catches wiring,
# gradient, NaN, or memory failures before the 50B run starts.
if resume_path is None:
run_startup_model_probe(model, len(tokenizer))
def create_optimizer(model):
return torch.optim.AdamW(
model.parameters(),
lr=PEAK_LR,
betas=(0.9, 0.95),
eps=1e-8,
weight_decay=WEIGHT_DECAY,
fused=True,
)
optimizer = create_optimizer(model)
progress = {
"step": 0,
"tokens": 0,
"sessions": 0,
"skipped_updates": 0,
"consecutive_skips": 0,
}
if PREFLIGHT_SMOKE:
def smoke_assert(condition, message):
passed = torch.tensor(int(condition), device=DEVICE, dtype=torch.int32)
dist.all_reduce(passed, op=dist.ReduceOp.MIN)
if not bool(passed.item()):
raise RuntimeError(message)
def state_fingerprint(value):
# Exact content check without retaining a second 2B model/Adam
# state in memory. Bound each device-to-host copy to 16 MiB.
digest = hashlib.sha256()
def update(item):
if torch.is_tensor(item):
digest.update(str((item.dtype, tuple(item.shape))).encode())
flat = item.detach().reshape(-1)
chunk_elements = max(1, (16 * 1024**2) // item.element_size())
for start in range(0, flat.numel(), chunk_elements):
chunk = flat[start:start + chunk_elements].to("cpu").contiguous()
digest.update(chunk.view(torch.uint8).numpy().tobytes())
elif isinstance(item, dict):
digest.update(b"dict")
for key in sorted(item, key=lambda key: (type(key).__name__, repr(key))):
update(key)
update(item[key])
elif isinstance(item, (list, tuple)):
digest.update(type(item).__name__.encode())
for child in item:
update(child)
else:
digest.update(repr(item).encode())
update(value)
return digest.hexdigest()
def matching_results(expected, actual):
if isinstance(expected, dict):
return isinstance(actual, dict) and expected.keys() == actual.keys() and all(
matching_results(value, actual[key]) for key, value in expected.items()
)
if isinstance(expected, (list, tuple)):
return isinstance(actual, (list, tuple)) and len(expected) == len(actual) and all(
matching_results(left, right) for left, right in zip(expected, actual)
)
if isinstance(expected, float):
return math.isfinite(actual) and math.isclose(
expected, actual, rel_tol=1e-6, abs_tol=1e-7
)
return expected == actual
generator = torch.Generator(device=DEVICE).manual_seed(12345 + RANK)
def synthetic_update(candidate, candidate_optimizer, step):
candidate.train()
candidate_optimizer.zero_grad(set_to_none=True)
candidate.module.set_specialization_scale(
min(1.0, step / max(1, SPEC_WARMUP_STEPS))
)
# Synchronize every microstep: rank-local top-1 expert usage can
# change between microsteps and ranks.
for _ in range(GRAD_ACCUM_STEPS):
x = torch.randint(
0, len(tokenizer), (MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device=DEVICE, dtype=torch.long, generator=generator,
)
y = torch.randint(
0, len(tokenizer), (MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device=DEVICE, dtype=torch.long, generator=generator,
)
with torch.autocast(
"cuda", dtype=torch.bfloat16, cache_enabled=AUTOCAST_CACHE
):
ce, auxiliary = candidate(x, y)
loss = (ce + auxiliary) / GRAD_ACCUM_STEPS
smoke_assert(
bool(torch.isfinite(loss.detach()).item()),
"Preflight smoke produced a nonfinite loss.",
)
loss.backward()
grads = gradient_health_summary(candidate)
smoke_assert(
all(math.isfinite(value) and value > 0 for value in grads.values()),
f"Preflight smoke produced invalid gradient groups: {grads}",
)
norm = torch.nn.utils.clip_grad_norm_(candidate.parameters(), GRAD_CLIP)
smoke_assert(
bool(torch.isfinite(norm).item()),
"Preflight smoke produced nonfinite gradients.",
)
for group in candidate_optimizer.param_groups:
group["lr"] = learning_rate(step)
candidate_optimizer.step()
candidate_optimizer.zero_grad(set_to_none=True)
return x, y
rank0_print(
f"Running synthetic optimizer update ({GRAD_ACCUM_STEPS} synchronized microsteps)...",
flush=True,
)
probe_x, probe_y = synthetic_update(model, optimizer, 1)
progress.update({
"step": 1,
"tokens": TOKENS_PER_UPDATE,
"sessions": 1,
})
expected_telemetry = run_telemetry_probe(model, probe_x, probe_y)
expected_ablation = run_ablation_probe(model, probe_x, probe_y)
smoke_assert(
all(math.isfinite(value) for value in expected_ablation.values())
and math.isfinite(expected_telemetry["probe_ce"])
and math.isfinite(expected_telemetry["probe_aux"]),
"Preflight smoke produced nonfinite probe results.",
)
expected_model = state_fingerprint(model.module.state_dict())
expected_optimizer = state_fingerprint(optimizer.state_dict())
gc.collect()
torch.cuda.empty_cache()
# The replicated DDP checkpoint format is exercised locally. The data
# state is deliberately absent because this path never opens the
# 50B-token stream.
smoke_data_state = {"mode": "preflight_smoke", "cursor": None}
checkpoint_writer = AsyncCheckpointWriter() if RANK == 0 else None
smoke_error = None
if RANK == 0:
try:
checkpoint_writer.submit(capture_local_checkpoint(
model, optimizer, tokenizer, progress, smoke_data_state
))
except Exception as exc:
smoke_error = f"Smoke checkpoint capture failed: {type(exc).__name__}: {exc}"
smoke_error = broadcast_object(smoke_error)
if smoke_error is not None:
raise RuntimeError(smoke_error)
# Deliberately mutate the live model while the captured checkpoint
# writes: restore must match the older fingerprints, not this update.
overlap_x, overlap_y = synthetic_update(model, optimizer, 2)
del overlap_x, overlap_y
smoke_result = None
if RANK == 0:
try:
smoke_result = {"path": str(checkpoint_writer.wait())}
except Exception as exc:
smoke_result = {"error": f"Smoke checkpoint write failed: {type(exc).__name__}: {exc}"}
smoke_result = broadcast_object(smoke_result)
if "error" in smoke_result:
raise RuntimeError(smoke_result["error"])
smoke_path = smoke_result["path"]
dist_barrier()
rank0_print("Local preflight checkpoint saved:", smoke_path, flush=True)
del model, optimizer
gc.collect()
torch.cuda.empty_cache()
fresh_model = allocate_model(False)
fresh_optimizer = create_optimizer(fresh_model)
restored = restore_checkpoint(
Path(smoke_path), fresh_model, fresh_optimizer
)
smoke_assert(
restored.get("progress") == progress
and restored.get("data_state") == smoke_data_state,
"Preflight smoke restore returned unexpected progress/data state.",
)
smoke_assert(
state_fingerprint(fresh_model.module.state_dict()) == expected_model,
"Preflight smoke restored model tensors differ from the saved model.",
)
smoke_assert(
state_fingerprint(fresh_optimizer.state_dict()) == expected_optimizer,
"Preflight smoke restored optimizer tensors/groups differ from the saved optimizer.",
)
fresh_model.module.set_specialization_scale(
min(1.0, 1.0 / max(1, SPEC_WARMUP_STEPS))
)
actual_telemetry = run_telemetry_probe(fresh_model, probe_x, probe_y)
actual_ablation = run_ablation_probe(fresh_model, probe_x, probe_y)
smoke_assert(
matching_results(expected_telemetry, actual_telemetry)
and matching_results(expected_ablation, actual_ablation),
"Preflight smoke restored numerical telemetry/ablation results differ.",
)
# A successful state load alone does not prove the fresh DDP reducer
# can perform another backward with dynamically unused experts.
del probe_x, probe_y
post_x, post_y = synthetic_update(fresh_model, fresh_optimizer, 2)
del post_x, post_y
smoke_assert(
state_fingerprint(fresh_model.module.state_dict()) != expected_model
and state_fingerprint(fresh_optimizer.state_dict()) != expected_optimizer,
"Preflight smoke post-restore update did not change model/optimizer state.",
)
dist_barrier()
rank0_print(
"PREFLIGHT SMOKE PASS: accumulated update, telemetry + ablation, "
"async local save with live model mutation, exact fresh model/optimizer restore, numerical "
"probe equivalence, and post-restore backward/update completed.",
flush=True,
)
del fresh_model, fresh_optimizer, restored
return
saved_data = None
if resume_path is not None:
rank0_print("Restoring 2B DDP model + optimizer...", flush=True)
state = restore_checkpoint(resume_path, model, optimizer)
progress.update(state["progress"])
saved_data = state.get("data_state")
del state
if progress["tokens"] > TRAIN_TOKENS or progress["step"] > MAX_STEPS:
raise RuntimeError("Checkpoint progress exceeds this frozen 2B run target.")
progress["sessions"] += 1
progress.setdefault("skipped_updates", 0)
progress.setdefault("consecutive_skips", 0)
if isinstance(saved_data, dict):
if int(saved_data.get("version", -1)) != DATA_STREAM_VERSION:
raise ValueError(
f"Checkpoint data stream v{saved_data.get('version')!r} is incompatible "
f"with stratified stream v{DATA_STREAM_VERSION}. Start a clean v3 run."
)
if saved_data.get("sequence_length") != SEQUENCE_LENGTH:
raise ValueError("Saved data sequence length differs.")
spec = saved_data.get("spec", {})
prefix = DATA_CACHE_DIR.rstrip("/") + "/"
files = list(spec.get("files", []))
if (
spec.get("version") != 2
or spec.get("kind") != "nano_base_hub"
or spec.get("repo_id") != DATA_REPO_ID
or spec.get("repo_type") != DATA_REPO_TYPE
or not files
or not all(str(name).startswith(prefix) for name in files)
):
raise ValueError("Saved data state is not this Nano-base corpus.")
if int(spec.get("vocab_size", -1)) != len(tokenizer):
raise ValueError("Saved data vocabulary differs.")
seed = int(saved_data["seed"])
cursor = saved_data["cursor"]
rank0_print(
"2B DDP RESUME: restoring exact data cursor:", cursor,
"\nPinned data revision:", spec["revision"], flush=True,
)
else:
spec = create_data_spec(len(tokenizer))
seed = DATA_SEED
cursor = None
rank0_print("2B DDP FRESH DATA STREAM: beginning at corpus cursor 0.", flush=True)
source = ShardStream(spec, seed, cursor) if RANK == 0 else None
committed_data = {
"version": DATA_STREAM_VERSION,
"sequence_length": SEQUENCE_LENGTH,
"seed": seed,
"spec": spec,
"cursor": source.state_dict() if RANK == 0 else cursor,
}
stream = PrefetchedStream(source) if RANK == 0 else None
if downloaded:
shutil.rmtree(WORK / "download", ignore_errors=True)
gc.collect()
torch.cuda.empty_cache()
for sig in (signal.SIGINT, signal.SIGTERM):
try:
old_handlers[sig] = signal.signal(sig, request_stop)
except ValueError:
pass
latest_local = resume_path if resume_path is not None and not downloaded else None
last_save_time = time.monotonic()
dirty = saved_data is None
in_optimizer_step = False
checkpoint_inflight = None
checkpoint_writer = AsyncCheckpointWriter() if RANK == 0 else None
checkpoint_size_estimate = (
total * 4
+ max(nested_bytes(optimizer.state), total * 8)
+ 512 * 1024**2
)
def train_deadline():
upload_estimate = (
hub.estimate_upload_seconds(checkpoint_size_estimate)
if hub is not None else 0.0
)
reserve = max(
UPLOAD_RESERVE_MINUTES * 60,
1.25 * upload_estimate + SAVE_RESERVE_MINUTES * 60,
)
return min(nominal_train_deadline, hard_deadline - reserve)
def finish_checkpoint(wait=False):
nonlocal latest_local, last_save_time, dirty, checkpoint_size_estimate
nonlocal checkpoint_inflight
if checkpoint_inflight is None:
return
saved = None
failure = None
if RANK == 0:
try:
# A periodic completion must not join an outstanding network
# upload. Keep the completed future until upload is free,
# but surface writer errors promptly even while uploading.
defer = not wait and hub.busy() and checkpoint_writer.failure() is None
path = None if defer else (
checkpoint_writer.wait() if wait else checkpoint_writer.result()
)
if path is not None:
if not checkpoint_valid(path):
raise RuntimeError("Local writer returned an incomplete checkpoint.")
size = folder_bytes(path)
# The old upload may still use its local files/staging.
# Join it before pruning or starting the next upload.
hub.wait()
prune_local_checkpoints(path)
rank0_print(
f"Local save finished in "
f"{time.monotonic() - checkpoint_inflight['start']:.1f}s | "
f"{size / 1024**3:.2f} GiB\n"
f"Local checkpoint: {path}", flush=True,
)
hub.upload_async(path, checkpoint_inflight["progress"])
saved = {
"path": str(path), "size": size, "time": time.monotonic(),
"progress": checkpoint_inflight["progress"],
}
except BaseException as exc:
failure = exc
saved = {"error": f"Checkpoint finalization failed: {type(exc).__name__}: {exc}"}
saved = broadcast_object(saved)
if saved is None:
return
if "error" in saved:
checkpoint_inflight = None
dirty = True
if failure is not None:
raise failure
raise RuntimeError(saved["error"])
latest_local = Path(saved["path"])
checkpoint_size_estimate = saved["size"]
last_save_time = saved["time"]
# Training may have advanced while this older snapshot was written.
dirty = progress != saved["progress"]
checkpoint_inflight = None
def save_checkpoint(reason):
nonlocal checkpoint_inflight
periodic = reason == "periodic"
if checkpoint_inflight is not None:
if periodic:
return
finish_checkpoint(wait=True)
if not dirty:
return
optimizer.zero_grad(set_to_none=True)
rank0_print(f"\nSaving ({reason}) at step {progress['step']:,}...", flush=True)
queued = None
failure = None
if RANK == 0:
try:
snapshot = capture_local_checkpoint(
model, optimizer, tokenizer, progress, committed_data
)
queued = {
"progress": copy.deepcopy(snapshot["progress"]),
"start": time.monotonic(),
}
checkpoint_writer.submit(snapshot)
except BaseException as exc:
failure = exc
queued = {"error": f"Checkpoint capture failed: {type(exc).__name__}: {exc}"}
queued = broadcast_object(queued)
if "error" in queued:
if failure is not None:
raise failure
raise RuntimeError(queued["error"])
checkpoint_inflight = queued
if not periodic:
finish_checkpoint(wait=True)
metrics = torch.zeros(2, device=DEVICE, dtype=torch.float32)
telemetry_probe = None
latest_grad_health = {k: 0.0 for k in (
"cortex", "procedure", "vault", "state", "embedding"
)}
log_updates = 0
log_tokens = 0
log_start = time.monotonic()
previous_retries = torch.cuda.memory_stats().get("num_alloc_retries", 0)
torch.cuda.reset_peak_memory_stats()
rank0_print(
f"Starting optimizer step {progress['step']:,}; "
f"{TRAIN_TOKENS - progress['tokens']:,} tokens remain.", flush=True,
)
try:
while progress["step"] < MAX_STEPS and progress["tokens"] < TRAIN_TOKENS:
# Surface writer failures to every rank before the next update.
finish_checkpoint()
stop_flag = torch.tensor(
int(RANK == 0 and (STOP_REQUESTED or time.monotonic() >= train_deadline())),
device=DEVICE,
)
dist.broadcast(stop_flag, src=0)
if bool(stop_flag.item()):
break
optimizer.zero_grad(set_to_none=True)
ce_sum = torch.zeros((), device=DEVICE)
aux_sum = torch.zeros((), device=DEVICE)
next_cursor = None
next_step = progress["step"] + 1
spec_scale = min(1.0, next_step / max(1, SPEC_WARMUP_STEPS))
model.module.set_specialization_scale(spec_scale)
for _ in range(GRAD_ACCUM_STEPS):
x, y, next_cursor = distributed_batch(stream)
try:
with torch.autocast(
"cuda", dtype=torch.bfloat16,
cache_enabled=AUTOCAST_CACHE,
):
ce, auxiliary = model(x, y)
loss = (ce + auxiliary) / GRAD_ACCUM_STEPS
except torch.cuda.OutOfMemoryError as exc:
raise RuntimeError(
"2B DDP profile OOM. Set GLOBAL_MICRO_BATCH_SIZE=8 "
"(1 sequence/GPU, grad_accum=14, 114,688 tokens/update) "
"before changing the architecture."
) from exc
telemetry_probe = (x.detach().clone(), y.detach().clone())
loss.backward()
ce_sum.add_(ce.detach())
aux_sum.add_(auxiliary.detach())
del x, y, ce, auxiliary, loss
if T2_TELEMETRY_EVERY_STEPS > 0 and next_step % T2_TELEMETRY_EVERY_STEPS == 0:
latest_grad_health = gradient_health_summary(model)
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)
finite = bool(torch.isfinite(norm).item())
finite_tensor = torch.tensor(int(finite), device=DEVICE)
dist.all_reduce(finite_tensor, op=dist.ReduceOp.MIN)
finite = bool(finite_tensor.item())
lr = learning_rate(next_step)
for group in optimizer.param_groups:
group["lr"] = lr
if finite:
in_optimizer_step = True
optimizer.step()
in_optimizer_step = False
progress["step"] += 1
progress["consecutive_skips"] = 0
metrics[0].add_(ce_sum / GRAD_ACCUM_STEPS)
metrics[1].add_(aux_sum / GRAD_ACCUM_STEPS)
log_updates += 1
else:
progress["skipped_updates"] += 1
progress["consecutive_skips"] += 1
rank0_print(
"Nonfinite gradient: skipped update. "
f"Consecutive={progress['consecutive_skips']}, "
f"lifetime={progress['skipped_updates']}", flush=True,
)
optimizer.zero_grad(set_to_none=True)
progress["tokens"] += TOKENS_PER_UPDATE
if RANK == 0:
committed_data["cursor"] = copy.deepcopy(next_cursor)
dirty = True
log_tokens += TOKENS_PER_UPDATE
del ce_sum, aux_sum, norm
if progress["consecutive_skips"] >= MAX_CONSECUTIVE_NONFINITE:
save_checkpoint("consecutive nonfinite gradients")
raise RuntimeError("Too many consecutive nonfinite updates.")
if finite and progress["step"] % LOG_EVERY_STEPS == 0:
values = metrics.cpu().tolist()
elapsed = time.monotonic() - log_start
allocated = torch.cuda.memory_allocated() / 1024**3
reserved = torch.cuda.memory_reserved() / 1024**3
peak = torch.cuda.max_memory_allocated() / 1024**3
retries = torch.cuda.memory_stats().get("num_alloc_retries", 0)
data_wait = stream.wait_seconds if RANK == 0 else 0.0
if RANK == 0:
stream.wait_seconds = 0.0
rank0_print(
f"step={progress['step']:,}/{MAX_STEPS:,} "
f"tokens={progress['tokens']:,}/{TRAIN_TOKENS:,} "
f"nominal_corpus_passes={progress['tokens'] / DATA_TARGET_TOKENS:.6f}/{TRAIN_EPOCHS} "
f"ce={values[0] / max(1, log_updates):.4f} "
f"aux={values[1] / max(1, log_updates):.4f} "
f"spec={spec_scale:.2f} lr={lr:.3e} "
f"tok/s={log_tokens / max(elapsed, 1e-6):,.0f} "
f"VRAM={allocated:.1f}/{reserved:.1f}GiB peak={peak:.1f}GiB "
f"data_wait={data_wait:.3f}s "
f"alloc_retries_delta={retries - previous_retries} "
f"skips={progress['skipped_updates']} "
f"session_left={max(0, train_deadline() - time.monotonic()) / 3600:.2f}h",
flush=True,
)
metrics.zero_()
log_updates = 0
log_tokens = 0
log_start = time.monotonic()
previous_retries = retries
torch.cuda.reset_peak_memory_stats()
if (
finite and telemetry_probe is not None
and T2_TELEMETRY_EVERY_STEPS > 0
and progress["step"] % T2_TELEMETRY_EVERY_STEPS == 0
):
probe_x, probe_y = telemetry_probe
t2 = run_telemetry_probe(model, probe_x, probe_y)
bank_text = " ".join(
f"B{b['bank']}:H={b['router_entropy']:.2f},"
f"load={100*b['min_load']:.1f}-{100*b['max_load']:.1f}%,"
f"dead={b['dead_experts']:.0f},null={b['null_probability']:.3f},"
f"shared={b['shared_gate']:.3f}"
for b in t2["banks"]
)
rank0_print(
"T2.1 "
f"step={progress['step']:,} probe_ce={t2['probe_ce']:.4f} "
f"vault_gate={t2['vault_gate']:.3f}±{t2['vault_gate_std']:.3f} "
f"vault_H={t2['vault_entropy']:.3f} "
f"vault_top1={t2['vault_top1']:.3f} "
f"vault_margin={t2['vault_margin']:.3f} "
f"vault_unique={t2['vault_unique']:.0f}/{model.module.vault.slots} "
f"state_read={t2['state_read']:.3f} "
f"state_write={t2['state_write']:.3f} "
f"state_delta={t2['state_delta']:.4f} state_rms={t2['state_rms']:.4f} "
f"grad[C/P/V/S]={latest_grad_health['cortex']:.2e}/"
f"{latest_grad_health['procedure']:.2e}/"
f"{latest_grad_health['vault']:.2e}/"
f"{latest_grad_health['state']:.2e} " + bank_text,
flush=True,
)
warnings = health_warnings(t2, latest_grad_health, progress["step"])
if warnings:
rank0_print("🚨 T2.1 HEALTH WARNINGS: " + " | ".join(warnings), flush=True)
if (
finite and telemetry_probe is not None
and T2_ABLATION_EVERY_STEPS > 0
and progress["step"] % T2_ABLATION_EVERY_STEPS == 0
):
probe_x, probe_y = telemetry_probe
n = min(ABLATION_PROBE_BATCH, probe_x.shape[0])
a = run_ablation_probe(model, probe_x[:n], probe_y[:n])
rank0_print(
"ABLATE "
f"step={progress['step']:,} full={a['full']:.4f} "
f"Δvault={a['no_vault']-a['full']:+.4f} "
f"Δstate={a['no_state']-a['full']:+.4f} "
f"Δprocedure={a['no_procedure']-a['full']:+.4f} "
f"Δdelib={a['no_deliberation']-a['full']:+.4f}",
flush=True,
)
# Only rank zero owns the upload state and scheduling clock.
# All ranks must nevertheless enter every collective save.
finish_checkpoint()
should_save = False
if RANK == 0:
now = time.monotonic()
upload_estimate = hub.estimate_upload_seconds(checkpoint_size_estimate)
final_guard = max(
FINAL_SAVE_GUARD_MINUTES * 60,
2 * upload_estimate + SAVE_RESERVE_MINUTES * 60,
)
should_save = (
dirty and not STOP_REQUESTED
and checkpoint_inflight is None and not checkpoint_writer.busy()
and not hub.busy()
and now - last_save_time >= CHECKPOINT_EVERY_MINUTES * 60
and train_deadline() - now > final_guard
)
if broadcast_object(should_save):
save_checkpoint("periodic")
metrics.zero_()
log_updates = 0
log_tokens = 0
log_start = time.monotonic()
had_checkpoint_inflight = checkpoint_inflight is not None
finish_checkpoint(wait=True)
if broadcast_object(dirty if RANK == 0 else None):
save_checkpoint("final")
elif RANK == 0 and latest_local is not None and not had_checkpoint_inflight:
hub.wait()
pointer = hub.pointer
remote_key = (
(int(pointer["step"]), int(pointer.get("tokens", 0)))
if pointer else (-1, -1)
)
local_key = (progress["step"], progress["tokens"])
if local_key > remote_key:
hub.upload_async(latest_local, progress)
done = progress["tokens"] >= TRAIN_TOKENS or progress["step"] >= MAX_STEPS
rank0_print(
f"Done session: step {progress['step']:,}, "
f"{progress['tokens']:,} consumed tokens.", flush=True,
)
if done:
rank0_print(f"{MODEL_NAME} BASE TARGET COMPLETE.", flush=True)
else:
remaining = max(0, TRAIN_TOKENS - progress["tokens"])
rank0_print(
f"Session ended normally; {remaining:,} tokens remain. "
"Rerun this exact script to restore model, optimizer, and data cursor.",
flush=True,
)
except DataPipelineError:
optimizer.zero_grad(set_to_none=True)
rank0_print(
"Data pipeline failed. Preserving the last complete optimizer boundary.",
flush=True,
)
if broadcast_object(
(dirty or checkpoint_inflight is not None) and not in_optimizer_step
if RANK == 0 else None
):
try:
save_checkpoint("data pipeline recovery")
except Exception as exc:
rank0_print("Recovery save failed:", exc)
raise
except torch.cuda.OutOfMemoryError:
optimizer.zero_grad(set_to_none=True)
gc.collect()
torch.cuda.empty_cache()
rank0_print(
"\nCUDA OOM. Existing checkpoints remain valid.\n"
"Set GLOBAL_MICRO_BATCH_SIZE=8 (1 sequence/GPU); "
"GRAD_ACCUM_STEPS becomes 14, keeping 114,688 tokens/update. "
"If already at 8, increase activation checkpointing; DDP cannot shard weights.\n"
"No checkpoint is taken from a partial optimizer update.",
flush=True,
)
raise
except Exception:
rank0_print(
"\nTraining stopped unexpectedly. Previously completed checkpoints remain valid.",
flush=True,
)
raise
finally:
if stream is not None:
stream.close()
for sig, handler in old_handlers.items():
try:
signal.signal(sig, handler)
except ValueError:
pass
try:
if checkpoint_writer is not None:
checkpoint_writer.close()
finally:
try:
if hub is not None:
rank0_print("Waiting for any pending checkpoint upload...", flush=True)
hub.close()
finally:
# Do not enter a new collective during exception unwinding: another
# rank may still be in a different collective when torchrun aborts it.
if dist.is_initialized():
dist.destroy_process_group()
if __name__ == "__main__":
main()