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