diff --git "a/App_ddp.py" "b/App_ddp.py" new file mode 100644--- /dev/null +++ "b/App_ddp.py" @@ -0,0 +1,4510 @@ +""" + +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": " 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()