""" Orion Flagship Nano — T2.1 single-file base-pretraining trainer. Target hardware: NVIDIA H200 (SM 90) or RTX PRO 6000 Blackwell (SM 120), ~96 GB VRAM. T2.1 changes over the original Mini prototype: - ~135M-parameter Nano configuration. - 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. It intentionally does NOT load Mini-T2 weights. Distributed launch (one node with eight supported GPUs): torchrun --standalone --nproc_per_node=8 App_fsdp.py The configured micro-batch is global (112 sequences); each rank processes 14 sequences, so the optimizer update still represents 114,688 tokens. This file uses classic PyTorch FSDP with a rank-zero full-state checkpoint format (v3). Format-2 single-GPU checkpoints are intentionally rejected rather than loaded with an unsafe optimizer conversion. """ import os import re import sys import importlib.util import subprocess # =========================================================================== # SETTINGS — edit these here, not in environment variables # ===========================================================================\ MODEL_NAME = "Orion Flagship Nano T2.1" # Dedicated rolling checkpoint repo for the Nano experiment. HF_REPO_ID = os.environ.get( "ORION_NANO_MODEL_REPO", "Project-Prism/Orion-Flagship-Nano-T2.1", ) 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_NANO_WORK_DIR", "./orion_nano_t21_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 # Large physical batch on the target ~96GB-class GPUs; base pretraining stays at # 1K context. SEQUENCE_LENGTH = 1024 GLOBAL_MICRO_BATCH_SIZE = 112 # For OOMs, set this to 56: 7 sequences/GPU × 8 GPUs × 2 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 # One almost-complete pass over the 50B-token cache; only the unavoidable final # partial optimizer batch is left unused rather than exceeding the requested budget. TRAIN_UPDATE_COUNT = DATA_TARGET_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 = 12 CHECKPOINT_DELIBERATION = False CHECKPOINT_LOSS_CHUNKS = False LOSS_CHUNK_TOKENS = 16_384 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 exercised by this run. 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-nano-t21-stratified-v3" 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_NANO_PREFLIGHT_SMOKE=1 torchrun --standalone --nproc_per_node=8 App_fsdp.py 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_NANO_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", "bitsandbytes": "bitsandbytes>=0.46", "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 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 bitsandbytes as bnb import triton import triton.language as tl from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, ShardingStrategy, StateDictType, FullStateDictConfig, FullOptimStateDictConfig, ) 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 FSDP 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 NANO — T2.1 ARCHITECTURE # =========================================================================== ARCH = { "model_name": MODEL_NAME, "implementation_version": 3, "architecture": "orion_flagship_nano_t2_1", "d_model": 640, "n_layers": 12, "n_heads": 10, "n_kv_heads": 2, "rope_theta": 10000.0, "sequence_length": SEQUENCE_LENGTH, "tokenizer_name": TOKENIZER_NAME, # Procedure Bank v2. Three banks are reused across four logical stages each. "n_procedure_banks": 3, "procedure_bank_span": 4, "n_experts": 12, "top_k": 1, "expert_hidden": 784, "shared_expert_hidden": 256, # Product-key Vault v2. Keep the 65,536-slot capacity but make the read # path state/stage-aware and confidence-gated instead of merely enlarging it. "vault_key_parts": 256, "vault_slots": 65_536, "vault_top_component": 12, "vault_top_k": 4, "vault_gate_groups": 64, "vault_read_layers": [3, 7, 11], "working_state_layers": [3, 7, 11], # One extra recurrent pass over the final three stages. "deliberation_start": 9, "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 OrionFlagshipNanoT21(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 FSDP wrapper.""" module = model.module if isinstance(model, FSDP) 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): """Run each ablation through the wrapper so FSDP gathers its parameters.""" 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_counts(cfg): """Exact meta count + approximate per-stage active compute proxy.""" with torch.device("meta"): probe = OrionFlagshipNanoT21(cfg) total = sum(p.numel() for p in probe.parameters()) 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): # Original parameters expose only this rank's shard under FULL_SHARD. # Keep every group in the collective even when this rank has no gradients. 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 = ".".join( part for part in name.split(".") if part != "_fsdp_wrapped_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 if dist.is_initialized(): dist.all_reduce(squared_norms, op=dist.ReduceOp.SUM) 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, ) model.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_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 manifest.get("run_id") != RUN_ID: return False files = ( manifest["model_shards"] + manifest["optimizer_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 save_local_checkpoint( model, optimizer, tokenizer, progress, data_state ): cfg = model.module.cfg if isinstance(model, FSDP) else model.cfg # FULL_STATE_DICT is collective: every rank must enter this context, but # rank zero is the only writer. This keeps the existing portable checkpoint # layout while making it valid for sharded parameters. state_cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) optim_cfg = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg): full_model = model.state_dict() full_optimizer = FSDP.optim_state_dict(model, optimizer) final = None if RANK == 0: try: final = _write_local_checkpoint( full_model, full_optimizer, tokenizer, progress, data_state, cfg ) except BaseException as exc: # Peers wait for this outcome, including failures before file writing # (directory creation, disk-space checks, CUDA synchronization). broadcast_object(f"{type(exc).__name__}: {exc}") raise error = broadcast_object(None) if error is not None: raise RuntimeError(f"Rank-zero checkpoint save failed: {error}") return final def _write_local_checkpoint( full_model, full_optimizer, tokenizer, progress, data_state, cfg ): 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) torch.cuda.synchronize() manifest = { "format_version": 3, "checkpoint_format": "fsdp_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(cpu_tree(shard), str(temporary / filename)) manifest["model_shards"].append(filename) optimizer_items = full_optimizer.items() for i, shard in enumerate(shard_items(optimizer_items)): filename = f"optimizer-{i:04d}.pt" torch.save(cpu_tree(shard), temporary / filename) manifest["optimizer_shards"].append(filename) state = { "progress": copy.deepcopy(progress), "optimizer_groups": None, "data_state": copy.deepcopy(data_state), "training_settings": training_settings(), "versions": { "torch": str(torch.__version__), "bitsandbytes": getattr(bnb, "__version__", "unknown"), }, **capture_rng(), } torch.save(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( "# Orion Flagship Nano T2.1 training checkpoint\n\n" "Custom PyTorch Orion Flagship Nano T2.1 architecture, not Transformers AutoModel.\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.") merged = {} 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.") duplicates = merged.keys() & shard.keys() if duplicates: raise ValueError( f"Duplicate keys in checkpoint {key} shard {filename}: " f"{sorted(map(str, duplicates))}" ) # shard_items splits top-level entries, including optimizer state and # param_groups; every listed shard contributes to the full dictionary. merged.update(shard) return merged def restore_checkpoint(path, model, optimizer): manifest = json.loads((path / "manifest.json").read_text()) if manifest.get("run_id") != RUN_ID: raise ValueError("Checkpoint belongs to an older/incompatible Nano run.") cfg = model.module.cfg if isinstance(model, FSDP) else model.cfg if manifest["architecture"] != cfg: raise ValueError("Checkpoint architecture does not match.") if manifest.get("checkpoint_format") != "fsdp_full_state": raise ValueError("This FSDP trainer only accepts format-3 checkpoints.") 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), ) # Every rank reads the complete state from the shared filesystem. Model # loading needs the full dictionary on each rank, and optimizer conversion # must shard each rank's dictionary rather than expect rank-zero-only input. state_cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=False) optim_cfg = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=False) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg): model.load_state_dict(model_state) load_optim = FSDP.optim_state_dict_to_load( model, optimizer, optimizer_state ) optimizer.load_state_dict(load_optim) state = torch.load( path / "training_state.pt", map_location="cpu", weights_only=False, ) 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_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." ) restore_rng(state) 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: return None required = ( list(manifest["model_shards"]) + list(manifest["optimizer_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 Nano 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("🏎️ ORION FLAGSHIP NANO — T2.1 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"Target: {TRAIN_TOKENS:,} tokens ({TRAIN_UPDATE_COUNT:,} updates); " f"{DATA_TARGET_TOKENS - TRAIN_TOKENS:,} corpus-tail tokens intentionally unused.", 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}." ) if not dist.is_initialized(): dist.init_process_group("nccl") torch.cuda.set_device(LOCAL_RANK) 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 NANO — bounded 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. hardware, inventory = validate_hardware() 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 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 132_000_000 <= total <= 138_000_000: raise RuntimeError( f"Nano parameter guard failed: expected ~135M, got {total:,}." ) def allocate_model(initialize_weights): rank0_print("Allocating FP32 T2.1 model...", flush=True) with torch.device("meta"): candidate = OrionFlagshipNanoT21(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 FSDP( candidate, sharding_strategy=ShardingStrategy.FULL_SHARD, use_orig_params=True, device_id=DEVICE, sync_module_states=False, limit_all_gathers=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: rank0_print("Running one bounded synthetic optimizer update...", flush=True) optimizer.zero_grad(set_to_none=True) x = torch.randint( 0, len(tokenizer), (MICRO_BATCH_SIZE, SEQUENCE_LENGTH), device=DEVICE, dtype=torch.long, ) y = torch.randint( 0, len(tokenizer), (MICRO_BATCH_SIZE, SEQUENCE_LENGTH), device=DEVICE, dtype=torch.long, ) model.set_specialization_scale( min(1.0, 1.0 / max(1, SPEC_WARMUP_STEPS)) ) with torch.autocast( "cuda", dtype=torch.bfloat16, cache_enabled=AUTOCAST_CACHE ): ce, auxiliary = model(x, y) loss = ce + auxiliary finite_loss = torch.isfinite(loss.detach()).to(dtype=torch.int32) dist.all_reduce(finite_loss, op=dist.ReduceOp.MIN) if not bool(finite_loss.item()): raise RuntimeError("Preflight smoke produced a nonfinite loss.") loss.backward() norm = model.clip_grad_norm_(GRAD_CLIP) finite_grad = torch.isfinite(norm).to(dtype=torch.int32) dist.all_reduce(finite_grad, op=dist.ReduceOp.MIN) if not bool(finite_grad.item()): raise RuntimeError("Preflight smoke produced nonfinite gradients.") lr = learning_rate(1) for group in optimizer.param_groups: group["lr"] = lr optimizer.step() optimizer.zero_grad(set_to_none=True) progress.update({ "step": 1, "tokens": TOKENS_PER_UPDATE, "sessions": 1, }) del x, y, ce, auxiliary, loss, norm gc.collect() torch.cuda.empty_cache() # The existing 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} smoke_path = save_local_checkpoint( model, optimizer, tokenizer, progress, smoke_data_state ) smoke_path = broadcast_object( str(smoke_path) if RANK == 0 else None ) 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 ) if ( restored.get("progress", {}).get("step") != 1 or restored.get("progress", {}).get("tokens") != TOKENS_PER_UPDATE ): raise RuntimeError("Preflight smoke restore returned unexpected progress.") dist_barrier() rank0_print( "PREFLIGHT SMOKE PASS: startup probe, one optimizer update, " "local save, and fresh model/optimizer restore completed.", flush=True, ) del fresh_model, fresh_optimizer, restored return saved_data = None if resume_path is not None: rank0_print("Restoring Nano 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 Nano 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( "NANO 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("NANO 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_size_estimate = ( total * 4 + max(nested_bytes(optimizer.state), total * 2) + 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 save_checkpoint(reason): nonlocal latest_local, last_save_time, dirty, checkpoint_size_estimate if RANK == 0: try: hub.wait() staging_root = WORK / "upload-staging" if staging_root.exists(): shutil.rmtree(staging_root) except BaseException as exc: broadcast_object(f"Checkpoint preparation failed: {type(exc).__name__}") raise preparation_error = broadcast_object(None) if preparation_error is not None: raise RuntimeError(preparation_error) optimizer.zero_grad(set_to_none=True) start = time.monotonic() rank0_print(f"\nSaving ({reason}) at step {progress['step']:,}...", flush=True) path = save_local_checkpoint( model, optimizer, tokenizer, progress, committed_data ) saved = None if RANK == 0: try: size = folder_bytes(path) prune_local_checkpoints(path) rank0_print( f"Local save finished in {time.monotonic() - start:.1f}s | " f"{size / 1024**3:.2f} GiB\n" f"Local checkpoint: {path}", flush=True, ) hub.upload_async(path, progress) saved = {"path": str(path), "size": size, "time": time.monotonic()} except BaseException as exc: broadcast_object({"error": f"Checkpoint finalization failed: {type(exc).__name__}"}) raise saved = broadcast_object(saved) if "error" in saved: raise RuntimeError(saved["error"]) latest_local = Path(saved["path"]) checkpoint_size_estimate = saved["size"] last_save_time = saved["time"] dirty = False 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: 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.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( "Nano turbo profile OOM. Set GLOBAL_MICRO_BATCH_SIZE=56 " "(7 sequences/GPU, grad_accum=2, 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 = model.clip_grad_norm_(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"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. 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 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() if broadcast_object(dirty if RANK == 0 else None): save_checkpoint("final") elif RANK == 0 and latest_local is not None: 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.wait() 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("🔥 ORION FLAGSHIP NANO T2.1 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 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=56 first (7 sequences/GPU); " "GRAD_ACCUM_STEPS becomes 2, keeping 114,688 tokens/update. " "If necessary try GLOBAL_MICRO_BATCH_SIZE=16 (2/GPU, accumulation=7).\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 if hub is not None: rank0_print("Waiting for any pending checkpoint upload...", flush=True) hub.close() # 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()