code-fdsp-v2 / App_fsdp.py
Bc-AI's picture
Upload App_fsdp.py
3c0a550 verified
Raw History Blame Contribute Delete
156 kB
"""
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": "<u2",
"uint32": "<u4",
"<u2": "<u2",
"<u4": "<u4",
}
dtype = dtypes.get(manifest.get("dtype"))
if dtype is None:
raise ValueError(f"Unsupported token dtype: {manifest.get('dtype')}")
if dtype == "<u2" and vocab_size > 65536:
raise ValueError("uint16 data cannot represent this vocabulary.")
tokenizer_name = manifest.get("tokenizer_name")
if tokenizer_name != TOKENIZER_NAME:
raise ValueError(
f"Data tokenizer {tokenizer_name!r} differs from {TOKENIZER_NAME!r}."
)
if int(manifest.get("vocab_size", -1)) != vocab_size:
raise ValueError("Data manifest vocabulary differs from model tokenizer.")
files = []
for worker in manifest.get("workers", {}).values():
for filename in worker.get("shard_files", []):
if not isinstance(filename, str):
raise ValueError("Expected filenames in manifest shard_files.")
filename = safe_relative_path(filename)
prefix = DATA_CACHE_DIR.rstrip("/") + "/"
if not filename.startswith(prefix):
filename = prefix + filename
files.append(filename)
files = sorted(set(files))
if not files:
raise ValueError("No Nano token shards are available in the manifest.")
print(
f"Nano dataset: {len(files)} shards pinned to revision {revision}.\n"
f"Manifest tokens: {written_tokens:,}; train budget: {TRAIN_TOKENS:,}.\n"
"Shard lengths are read from actual files; v3 training interleaves sources every sequence.",
flush=True,
)
return {
"version": 2,
"kind": "nano_base_hub",
"repo_id": DATA_REPO_ID,
"repo_type": DATA_REPO_TYPE,
"revision": revision,
"dtype": dtype,
"files": files,
"vocab_size": vocab_size,
"manifest_target_tokens": target_tokens,
"source_targets": expected_targets,
}
# ===========================================================================
# Stratified multi-shard data stream (v3)
# ===========================================================================
class ShardStream:
"""
Deterministic source-stratified interleaving across the whole Nano corpus.
Why this exists:
v2 shuffled shard order, but then consumed one entire ~268M-token shard
before moving on. Because shards are source-homogeneous, that created very
long single-domain runs (for example, hundreds of millions of Python
tokens in a row).
v3 fixes that without rewriting the 50B-token cache:
* Source selection follows an exact 25-sequence cycle:
15 FineWeb-Edu, 3 FineMath-4+, 2 OpenWebMath, 5 Code.
* The 25 source labels are independently shuffled every cycle.
* Within each source, shards are shuffled per source-epoch.
* Within each shard, 1024-token sequences are shuffled.
* Up to four source shards are naturally active at once, with an LRU
memmap cache and asynchronous next-shard downloads.
* The complete per-source cursor + global sequence index is checkpointed,
making resume exact and independent of batch boundaries.
No training sequence crosses a shard boundary. Tiny shard tails are omitted.
"""
SOURCE_NAMES = (
"general_fineweb_edu",
"math_finemath_4plus",
"math_openwebmath",
"code_python_clean",
)
def __init__(self, spec, seed, cursor=None):
self.spec = copy.deepcopy(spec)
self.seed = int(seed)
self.files = list(spec["files"])
self.dtype = np.dtype(spec["dtype"])
if spec.get("version") != 2 or spec.get("kind") != "nano_base_hub":
raise ValueError("This script expects a version-2 Nano-base Hub data spec.")
if not self.files:
raise ValueError("Empty data file list.")
# Split the immutable file list by the source id embedded by the
# pretokenizer in each shard filename.
self.source_files = {name: [] for name in self.SOURCE_NAMES}
for filename in self.files:
matches = [name for name in self.SOURCE_NAMES if f"-{name}.bin" in filename]
if len(matches) != 1:
raise ValueError(
f"Could not uniquely identify Nano source for shard: {filename}"
)
self.source_files[matches[0]].append(filename)
for source, names in self.source_files.items():
if not names:
raise ValueError(f"No shards found for required source {source!r}.")
names.sort()
# global_sequence chooses the source deterministically. Each source has
# an independent cursor through its own shuffled shards/sequences.
self.cursor = {
"stream_version": DATA_STREAM_VERSION,
"global_sequence": 0,
"sources": {
source: {
"epoch": 0,
"shard_position": 0,
"sequence_position": 0,
}
for source in self.SOURCE_NAMES
},
}
if cursor is not None:
if int(cursor.get("stream_version", -1)) != DATA_STREAM_VERSION:
raise ValueError(
f"Saved data stream version {cursor.get('stream_version')!r} "
f"is incompatible with v{DATA_STREAM_VERSION} stratified mixing."
)
self.cursor = copy.deepcopy(cursor)
if int(self.cursor["global_sequence"]) < 0:
raise ValueError("Negative global data cursor.")
for source in self.SOURCE_NAMES:
c = self.cursor["sources"].get(source)
if not isinstance(c, dict):
raise ValueError(f"Missing saved cursor for source {source!r}.")
if any(int(c[k]) < 0 for k in ("epoch", "shard_position", "sequence_position")):
raise ValueError(f"Negative data cursor for source {source!r}.")
if int(c["shard_position"]) >= len(self.source_files[source]):
raise ValueError(f"Invalid shard position for source {source!r}.")
self.cache_dir = WORK / "shard_cache"
self.cache_dir.mkdir(parents=True, exist_ok=True)
self._mapped = OrderedDict()
self._active = {}
self._shard_orders = {}
self._sequence_orders = {}
self._download_pool = ThreadPoolExecutor(
max_workers=DOWNLOAD_WORKERS, thread_name_prefix="shard-download"
)
self._futures = {}
def state_dict(self):
return copy.deepcopy(self.cursor)
def _source_cycle(self, cycle):
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, int(cycle), 303])
)
order = np.asarray(MIX_CYCLE, dtype=object)
return order[rng.permutation(len(order))]
def _source_for_global(self, global_sequence):
cycle, offset = divmod(int(global_sequence), len(MIX_CYCLE))
return str(self._source_cycle(cycle)[offset])
def _shard_order(self, source, epoch):
key = (source, int(epoch))
if key not in self._shard_orders:
source_id = self.SOURCE_NAMES.index(source)
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, source_id, int(epoch), 101])
)
self._shard_orders[key] = rng.permutation(len(self.source_files[source]))
return self._shard_orders[key]
def _download(self, name):
return hf_hub_download(
repo_id=self.spec["repo_id"],
repo_type=self.spec["repo_type"],
revision=self.spec["revision"],
filename=name,
token=HF_TOKEN,
cache_dir=str(self.cache_dir),
)
def _future_for(self, name):
future = self._futures.get(name)
if future is None:
future = self._download_pool.submit(self._download, name)
self._futures[name] = future
return future
def _open(self, name):
if name in self._mapped:
self._mapped.move_to_end(name)
return self._mapped[name]
future = self._futures.pop(name, None)
path = future.result() if future is not None else self._download(name)
byte_count = Path(path).stat().st_size
if byte_count % self.dtype.itemsize:
raise ValueError(f"Malformed shard byte length: {name}")
token_count = byte_count // self.dtype.itemsize
data = (
None
if token_count < SEQUENCE_LENGTH + 1
else np.memmap(path, mode="r", dtype=self.dtype)
)
self._mapped[name] = (data, token_count)
self._mapped.move_to_end(name)
while len(self._mapped) > MAPPED_SHARD_CACHE:
self._mapped.popitem(last=False)
print(
f"Shard ready: {name} | {token_count:,} actual tokens",
flush=True,
)
return data, token_count
def _schedule_next_for_source(self, source):
c = self.cursor["sources"][source]
epoch = int(c["epoch"])
position = int(c["shard_position"]) + 1
if position == len(self.source_files[source]):
position = 0
epoch += 1
index = int(self._shard_order(source, epoch)[position])
name = self.source_files[source][index]
if name not in self._mapped and name not in self._futures:
self._future_for(name)
def _advance_shard(self, source):
c = self.cursor["sources"][source]
c["sequence_position"] = 0
c["shard_position"] += 1
if c["shard_position"] == len(self.source_files[source]):
c["shard_position"] = 0
c["epoch"] += 1
self._active.pop(source, None)
def _activate(self, source):
# At most one source-epoch worth of shards looking for a usable shard.
for _ in range(len(self.source_files[source]) + 1):
c = self.cursor["sources"][source]
epoch = int(c["epoch"])
position = int(c["shard_position"])
key = (source, epoch, position)
active = self._active.get(source)
if active is not None and active["key"] == key:
return active
index = int(self._shard_order(source, epoch)[position])
name = self.source_files[source][index]
data, token_count = self._open(name)
sequence_count = max(0, (token_count - 1) // SEQUENCE_LENGTH)
if sequence_count == 0:
if c["sequence_position"] != 0:
raise ValueError("Saved cursor points into a tiny shard.")
self._advance_shard(source)
continue
if c["sequence_position"] > sequence_count:
raise ValueError(
f"Saved cursor exceeds actual shard length for source {source!r}."
)
source_id = self.SOURCE_NAMES.index(source)
rng = np.random.default_rng(
np.random.SeedSequence([self.seed, source_id, epoch, index, 202])
)
sequence_order = rng.permutation(sequence_count)
active = {
"key": key,
"name": name,
"data": data,
"sequence_order": sequence_order,
}
self._active[source] = active
self._schedule_next_for_source(source)
return active
raise ValueError(f"No shards for source {source!r} contain a complete sequence.")
def _copy_one(self, source, destination):
while True:
active = self._activate(source)
c = self.cursor["sources"][source]
position = int(c["sequence_position"])
order = active["sequence_order"]
if position < len(order):
break
self._advance_shard(source)
sample = int(order[position])
start = sample * SEQUENCE_LENGTH
values = active["data"][start:start + SEQUENCE_LENGTH + 1]
if len(values) != SEQUENCE_LENGTH + 1:
raise RuntimeError("Internal shard bounds invariant failed.")
np.copyto(destination, values, casting="safe")
c["sequence_position"] += 1
def batch(self):
storage = torch.empty(
(GLOBAL_MICRO_BATCH_SIZE, SEQUENCE_LENGTH + 1),
dtype=torch.long,
pin_memory=True,
)
array = storage.numpy()
counts = {name: 0 for name in self.SOURCE_NAMES}
for row in range(GLOBAL_MICRO_BATCH_SIZE):
global_sequence = int(self.cursor["global_sequence"])
source = self._source_for_global(global_sequence)
self._copy_one(source, array[row])
self.cursor["global_sequence"] = global_sequence + 1
counts[source] += 1
maximum = int(array.max())
if maximum >= int(self.spec["vocab_size"]):
raise ValueError(
f"Token ID {maximum} is outside vocabulary {self.spec['vocab_size']}."
)
return storage, self.state_dict()
def close(self):
self._download_pool.shutdown(wait=False, cancel_futures=True)
self._active.clear()
self._mapped.clear()
self._futures.clear()
class DataPipelineError(RuntimeError):
pass
class PrefetchedStream:
def __init__(self, source):
self.source = source
self.queue = queue.Queue(maxsize=max(1, PREFETCH_BATCHES))
self.stopped = threading.Event()
self.wait_seconds = 0.0
self.thread = threading.Thread(
target=self._worker,
name="token-prefetch",
daemon=True,
)
self.thread.start()
def _put(self, item):
while not self.stopped.is_set():
try:
self.queue.put(item, timeout=0.2)
return True
except queue.Full:
pass
return False
def _worker(self):
try:
while not self.stopped.is_set():
storage, cursor = self.source.batch()
if not self._put((True, storage, cursor)):
return
except Exception as exc:
self._put((False, exc, None))
finally:
self.source.close()
def batch(self):
started = time.monotonic()
while True:
try:
ok, value, cursor = self.queue.get(timeout=0.5)
break
except queue.Empty:
if not self.thread.is_alive():
raise DataPipelineError("Data producer exited.")
self.wait_seconds += time.monotonic() - started
if not ok:
raise DataPipelineError("Data producer failed") from value
# One contiguous H2D copy rather than two overlapping transfers.
storage = value.to("cuda", non_blocking=True)
return storage[:, :-1], storage[:, 1:], cursor
def close(self):
self.stopped.set()
self.thread.join(timeout=2.0)
def distributed_batch(stream):
"""Fetch one global batch on rank zero and give every rank a disjoint slice."""
local_batch = GLOBAL_MICRO_BATCH_SIZE // WORLD_SIZE
failure = None
status = None
if RANK == 0:
try:
x, y, cursor = stream.batch()
packed = torch.stack((x, y), dim=0)
status = "ok"
except Exception as exc:
failure = exc
status = "oom" if isinstance(exc, torch.cuda.OutOfMemoryError) else "failed"
# Never leave peers waiting for a batch when its producer has failed. All
# ranks take the same recovery path, using the last committed data cursor.
status = broadcast_object(status)
if status == "oom":
raise torch.cuda.OutOfMemoryError("Rank-zero batch preparation ran out of memory.") from failure
if status != "ok":
raise DataPipelineError("Rank-zero batch preparation failed.") from failure
if RANK != 0:
packed = torch.empty(
(2, GLOBAL_MICRO_BATCH_SIZE, SEQUENCE_LENGTH),
device=DEVICE,
dtype=torch.long,
)
cursor = None
if dist.is_initialized():
dist.broadcast(packed, src=0)
start = RANK * local_batch
end = start + local_batch
return packed[0, start:end], packed[1, start:end], cursor
# ===========================================================================
# Checkpoint storage
# ===========================================================================
def optimizer_layout(optimizer, model):
names = {id(p): name for name, p in model.named_parameters()}
layout = []
for group in optimizer.param_groups:
saved = {k: v for k, v in group.items() if k != "params"}
saved["param_names"] = [names[id(p)] for p in group["params"]]
layout.append(saved)
return layout
def training_settings():
return {
"peak_lr": PEAK_LR,
"min_lr": MIN_LR,
"warmup_steps": WARMUP_STEPS,
"max_steps": MAX_STEPS,
"weight_decay": WEIGHT_DECAY,
"grad_clip": GRAD_CLIP,
"router_aux_coef": ROUTER_AUX_COEF,
"router_z_coef": ROUTER_Z_COEF,
"router_specialization_coef": ROUTER_SPECIALIZATION_COEF,
"spec_warmup_steps": SPEC_WARMUP_STEPS,
"train_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()