""" 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()