#!/usr/bin/env python3 """ Nova-10B-Dense · HuggingFace Jobs · 4-8× NVIDIA H200 (Hopper, 141GB HBM3e) ~10B fully-dense transformer (no MoE, no routing, no MoD) Architecture: D1 Fused QKV – single Linear for Q+K+V, split after projection D2 Differential Attn – dual-stream diff attn + QK-Norm + head-gate (GQA/RoPE) D3 SwiGLU FFN – dense gated MLP, no routing D4 RMSNorm – pre-norm on attn + FFN D5 LayerScale – per-layer learnable residual scale (0.1 init) D6 RoPE – extended-theta (500k) rotary embeddings D7 Logit soft-cap – tanh(x/cap)*cap on final logits D8 FP8 – torchao rowwise (Hopper sm_90 / Blackwell sm_100+) Training stack: O1 Muon – Newton-Schulz orthogonalized momentum for 2-D weights O2 AdamW-8bit – bitsandbytes, for embed/norms/1-D params O3 FSDP2 – PyTorch FullyShardedDataParallel O4 Cosine LR – separate schedules for Muon vs AdamW D9 Balanced streaming – equal tokens per domain, new shards auto-detected D10 ctx-warmup – sequence-length curriculum Fixes applied vs original: FIX1 _has_fp8_support() checks major >= 9 (Hopper sm_90 = H100/H200) old _is_blackwell() checked >= 10, silently disabled fp8 on H200 FIX2 build_optimizers() called BEFORE wrap_fsdp() on the raw CPU model after FSDP wrap, named_parameters() walks flat-param structure and the 2D-weight classifier sees nothing -> empty muon_p -> ValueError FIX3 sync_module_states=False + identical CPU seeding replaces the giant NCCL broadcast that crashed on both Blackwell and early Hopper runs Launch: torchrun --nproc_per_node=4 nova10b_dense.py torchrun --nproc_per_node=8 nova10b_dense.py Debug: NOVA_DEBUG=1 torchrun --nproc_per_node=8 nova10b_dense.py NCCL safe fallback: NOVA_NCCL_SAFE=1 torchrun --nproc_per_node=8 nova10b_dense.py """ # ══════════════════════════════════════════════════════════════════════════════ # ENV SETUP — MUST HAPPEN BEFORE `import torch` # ══════════════════════════════════════════════════════════════════════════════ import os os.environ.setdefault("CUDA_DEVICE_ORDER", "PCI_BUS_ID") _DEBUG = os.environ.get("NOVA_DEBUG", "0") == "1" if _DEBUG: os.environ.setdefault("CUDA_LAUNCH_BLOCKING", "1") os.environ.setdefault("TORCH_USE_CUDA_DSA", "1") os.environ.setdefault("TORCH_SHOW_CPP_STACKTRACES", "1") os.environ.setdefault("NCCL_DEBUG", "INFO") os.environ.setdefault("NCCL_DEBUG_SUBSYS", "INIT,COLL") _NCCL_SAFE = os.environ.get("NOVA_NCCL_SAFE", "0") == "1" if _NCCL_SAFE: os.environ.setdefault("NCCL_P2P_DISABLE", "1") os.environ.setdefault("NCCL_SHM_DISABLE", "1") os.environ.setdefault("NCCL_ALGO", "Ring") # ── stdlib ────────────────────────────────────────────────────────────────── import gc import re import json import time import math import random import logging import itertools import functools import socket import traceback from dataclasses import dataclass, asdict, field from typing import Optional, List, Iterator # ── third-party ───────────────────────────────────────────────────────────── import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import torch.distributed as dist from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, MixedPrecision, ShardingStrategy, BackwardPrefetch, CPUOffload, FullStateDictConfig, StateDictType, ) from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from torch.utils.data import IterableDataset, DataLoader from huggingface_hub import HfApi, hf_hub_download, upload_file, list_repo_tree from huggingface_hub import login as hf_login try: import bitsandbytes as bnb HAS_BNB = True except ImportError: HAS_BNB = False try: from safetensors.torch import save_file as st_save HAS_ST = True except ImportError: HAS_ST = False try: from torchao.float8 import convert_to_float8_training, Float8LinearConfig HAS_TORCHAO = True except ImportError as _e: HAS_TORCHAO = False import sys print(f"[torchao import failed]: {_e}", file=sys.stderr, flush=True) # ── global env ────────────────────────────────────────────────────────────── os.environ["TOKENIZERS_PARALLELISM"] = "false" torch.set_float32_matmul_precision("high") torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True torch.backends.cudnn.benchmark = True # ── logging ────────────────────────────────────────────────────────────────── class _RankFilter(logging.Filter): def filter(self, record): record.rank = dist.get_rank() if dist.is_initialized() else 0 return True _rank_filter = _RankFilter() logging.basicConfig( level=logging.INFO, format="%(asctime)s | R%(rank)s | %(message)s", datefmt="%H:%M:%S", ) for _h in logging.root.handlers: _h.addFilter(_rank_filter) log = logging.getLogger("nova10b") log.addFilter(_rank_filter) # ══════════════════════════════════════════════════════════════════════════════ # DISTRIBUTED HELPERS # ══════════════════════════════════════════════════════════════════════════════ def setup_dist() -> int: dist.init_process_group(backend="nccl") local_rank = int(os.environ.get("LOCAL_RANK", 0)) torch.cuda.set_device(local_rank) return local_rank def cleanup_dist(): dist.destroy_process_group() def is_main() -> bool: return rank() == 0 def rank() -> int: return dist.get_rank() if dist.is_initialized() else 0 def world_size() -> int: return dist.get_world_size() if dist.is_initialized() else 1 def barrier(): if dist.is_initialized(): dist.barrier() def all_reduce_mean(tensor: torch.Tensor) -> torch.Tensor: dist.all_reduce(tensor, op=dist.ReduceOp.SUM) tensor.div_(world_size()) return tensor # ══════════════════════════════════════════════════════════════════════════════ # SEEDING # ══════════════════════════════════════════════════════════════════════════════ def seed_for_model_init(seed: int): """ Seeds CPU RNG identically on EVERY rank for model construction. Rank-INDEPENDENT by design — guarantees bit-identical model init across all ranks without needing FSDP sync_module_states broadcast. """ random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) # CPU generator only — no CUDA touch def seed_for_runtime(seed: int, device: torch.device): """ Rank-dependent seeding for everything AFTER model construction: data shuffling, dropout, sampling, etc. """ random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): with torch.cuda.device(device): torch.cuda.manual_seed(seed) # ══════════════════════════════════════════════════════════════════════════════ # STAGED STARTUP DIAGNOSTICS # ══════════════════════════════════════════════════════════════════════════════ def _stage(name: str): """ Context manager: logs entry, syncs CUDA before AND after, barriers all ranks, on exception logs rank/device/stage/traceback before re-raising. """ class _StageCtx: def __enter__(self): self.t0 = time.time() dev = torch.cuda.current_device() if torch.cuda.is_available() else -1 log.info( f"[stage:{name}] enter | rank={rank()} device=cuda:{dev} " f"host={socket.gethostname()}" ) if torch.cuda.is_available(): torch.cuda.synchronize() return self def __exit__(self, exc_type, exc_val, exc_tb): if exc_type is not None: dev = torch.cuda.current_device() if torch.cuda.is_available() else -1 log.error( f"[stage:{name}] FAILED | rank={rank()} device=cuda:{dev} | " f"{exc_type.__name__}: {exc_val}" ) log.error( "".join(traceback.format_exception(exc_type, exc_val, exc_tb)) ) return False if torch.cuda.is_available(): torch.cuda.synchronize() dt = time.time() - self.t0 log.info(f"[stage:{name}] ok ({dt:.2f}s) | rank={rank()}") barrier() return False return _StageCtx() def log_cuda_environment(local_rank: int): dev = local_rank try: props = torch.cuda.get_device_properties(dev) cap_major, cap_minor = torch.cuda.get_device_capability(dev) nccl_ver = ( torch.cuda.nccl.version() if hasattr(torch.cuda, "nccl") else "unknown" ) drv = ( torch.cuda.driver_version() if hasattr(torch.cuda, "driver_version") else "n/a" ) log.info( f"[env] rank={rank()} local_rank={local_rank} " f"gpu={props.name} sm_{cap_major}{cap_minor} " f"mem={props.total_memory/1024**3:.0f}GB " f"torch={torch.__version__} cuda={torch.version.cuda} " f"nccl={nccl_ver} driver={drv} " f"device_order={os.environ.get('CUDA_DEVICE_ORDER','default')} " f"nccl_safe={_NCCL_SAFE}" ) except Exception as e: log.warning(f"[env] rank={rank()} could not fully query CUDA env: {e}") def log_p2p_matrix(local_rank: int, ws: int): """ Logs pairwise P2P access capability. H200 SXM nodes have NVLink + NVSwitch so expect all-Y. Any N on H200 is a red flag — you are not getting the fabric you are paying for, escalate to infra. """ if not is_main() or not torch.cuda.is_available(): return try: n = torch.cuda.device_count() rows = [] for i in range(n): row = [] for j in range(n): if i == j: row.append("-") else: can = torch.cuda.can_device_access_peer(i, j) row.append("Y" if can else "N") rows.append(" ".join(row)) log.info("[p2p] pairwise CUDA P2P access matrix (expect all-Y on H200 NVLink):") for i, r in enumerate(rows): log.info(f"[p2p] gpu{i}: {r}") except Exception as e: log.warning(f"[p2p] matrix query failed: {e}") def cuda_health_check(device: torch.device): """Minimal per-rank CUDA sanity check: basic op + small matmul.""" try: x = torch.ones(4096, device=device, dtype=torch.bfloat16) y = torch.ones(4096, device=device, dtype=torch.bfloat16) z = x + y m = torch.randn(1024, 1024, device=device, dtype=torch.bfloat16) n = torch.randn(1024, 1024, device=device, dtype=torch.bfloat16) p = m @ n torch.cuda.synchronize(device) checksum = float(z.sum().item()) + float(p.sum().item()) log.info( f"[health] rank={rank()} device={device} basic op OK " f"(checksum={checksum:.1f})" ) except Exception as e: log.error( f"[health] rank={rank()} device={device} FAILED basic CUDA op: {e}" ) raise def nccl_collective_selftest(device: torch.device, ws: int): """ Isolated NCCL collective stress test run BEFORE FSDP wrap or any model code. Tests the exact same primitive classes (all_gather, broadcast) that FSDP uses internally. If this fails -> fault is environmental (NCCL/driver), not model code. If this passes -> fault is specific to FSDP/model collective patterns. On H200 SXM with NVLink this should pass all sizes quickly. """ sizes_mb = [1, 16, 64, 256, 1024, 4096] for mb in sizes_mb: n_elem = (mb * 1024 * 1024) // 2 # bf16 = 2 bytes/elem try: local = torch.full( (n_elem,), float(rank()), device=device, dtype=torch.bfloat16 ) gathered = [torch.zeros_like(local) for _ in range(ws)] dist.all_gather(gathered, local) torch.cuda.synchronize(device) bcast = local.clone() dist.broadcast(bcast, src=0) torch.cuda.synchronize(device) expected_sum = sum(range(ws)) * n_elem actual_sum = sum(float(g.sum().item()) for g in gathered) ok = abs(actual_sum - expected_sum) < max(1.0, expected_sum * 1e-3) log.info( f"[nccl_selftest] rank={rank()} size={mb}MB " f"all_gather+broadcast OK match={ok}" ) del local, gathered, bcast torch.cuda.empty_cache() except Exception as e: log.error( f"[nccl_selftest] rank={rank()} size={mb}MB FAILED: {e}\n" f"On H200 SXM this is unexpected — check:\n" f" nvidia-smi topo -m (NVLink topology)\n" f" dmesg -T | grep Xid (GPU hardware errors)\n" f" NOVA_NCCL_SAFE=1 (disable P2P fast paths)\n" f"Escalate to infra if raw nccl-tests also fail." ) raise if is_main(): log.info( f"[nccl_selftest] ALL sizes up to {sizes_mb[-1]}MB passed " f"on all {ws} ranks — NCCL/NVLink healthy ✓" ) # ══════════════════════════════════════════════════════════════════════════════ # HF HUB # ══════════════════════════════════════════════════════════════════════════════ HF_TOKEN = os.environ.get("HF_TOKEN", "") HF_DATASET_REPO = os.environ.get("HF_DATASET_REPO", "ml-intern-explorers/nova1-pretrain-20T") HF_MODEL_REPO = os.environ.get("HF_MODEL_REPO", "ml-intern-explorers/nova-10b-dense-ckpts") DOMAINS = os.environ.get("DOMAINS", "code,general,math,reasoning").split(",") SHARD_PREFIX = "data" LOCAL_CACHE = "/tmp/nova10b_cache" TEMP_CKPT_DIR = "/tmp/nova10b_ckpts" for _d in [LOCAL_CACHE, TEMP_CKPT_DIR]: os.makedirs(_d, exist_ok=True) hf_api: Optional[HfApi] = None def init_hf(): global hf_api if not HF_TOKEN: raise RuntimeError( "HF_TOKEN not set.\n" " HF Jobs: Settings -> Secrets -> add HF_TOKEN\n" " Local: export HF_TOKEN=hf_xxx" ) hf_login(token=HF_TOKEN, add_to_git_credential=False) hf_api = HfApi() if is_main(): log.info("[hf] login ok") def hf_upload(local: str, remote: str) -> bool: try: upload_file( path_or_fileobj=local, path_in_repo=remote, repo_id=HF_MODEL_REPO, repo_type="model", token=HF_TOKEN, ) if is_main(): log.info(f"[hf] ↑ {remote}") return True except Exception as e: if is_main(): log.error(f"[hf] upload failed {remote}: {e}") return False def hf_download_dataset(path: str) -> str: return hf_hub_download( repo_id=HF_DATASET_REPO, filename=path, repo_type="dataset", cache_dir=LOCAL_CACHE, token=HF_TOKEN, ) def hf_download_ckpt(path: str) -> Optional[str]: try: return hf_hub_download( repo_id=HF_MODEL_REPO, filename=path, repo_type="model", cache_dir=TEMP_CKPT_DIR, token=HF_TOKEN, ) except Exception as e: if is_main(): log.error(f"[hf] ckpt download failed {path}: {e}") return None def list_shards(domain: str) -> List[str]: try: items = list_repo_tree( HF_DATASET_REPO, repo_type="dataset", path_in_repo=f"{SHARD_PREFIX}/{domain}", recursive=False, token=HF_TOKEN, ) return sorted( i.path for i in items if hasattr(i, "path") and i.path.endswith(".npy") ) except Exception as e: if is_main(): log.warning(f"[data] list_shards({domain}) error: {e}") return [] def _ckpt_step(fname: str) -> int: m = re.search(r"_s(\d+)\.", fname) return int(m.group(1)) if m else 0 def rotate_ckpts(tag: str, keep: int = 2): if not is_main(): return try: files = hf_api.list_repo_files( repo_id=HF_MODEL_REPO, repo_type="model", token=HF_TOKEN, ) tagged = sorted( [f for f in files if tag in f and f.endswith(".model.pt")], key=_ckpt_step, ) while len(tagged) > keep: victim = tagged.pop(0) for fn in [victim, victim.replace(".model.pt", ".meta.pt")]: try: hf_api.delete_file( path_in_repo=fn, repo_id=HF_MODEL_REPO, repo_type="model", token=HF_TOKEN, ) except Exception: pass log.info(f"[hf] rotated {victim}") except Exception as e: log.warning(f"[hf] rotate failed: {e}") # ══════════════════════════════════════════════════════════════════════════════ # CONFIG — ~10B dense # ══════════════════════════════════════════════════════════════════════════════ @dataclass class NovaConfig: vocab_size: int = 65536 d_model: int = 4096 n_heads: int = 32 n_kv_heads: int = 8 n_layers: int = 35 max_len: int = 2048 ffn_hidden: int = 14336 diff_lambda_init: float = 0.8 logit_softcap: float = 50.0 rope_theta: float = 500_000.0 use_fp8: str = "auto" lr: float = 2e-4 min_lr: float = 2e-5 muon_lr: float = 6e-3 muon_min_lr: float = 6e-4 weight_decay: float = 0.1 beta1: float = 0.9 beta2: float = 0.95 muon_momentum: float = 0.95 grad_clip: float = 1.0 warmup_steps: int = 4_000 total_steps: int = 130_000 ctx_warmup_steps: int = 3_000 ctx_start_len: int = 2048 ctx_end_len: int = 2048 micro_batch: int = 6 grad_accum: int = 8 log_every: int = 10 ckpt_every: int = 1_000 eval_every: int = 500 ckpt_keep: int = 2 seed: int = 42 max_hours: float = 47.5 domains: List[str] = field(default_factory=lambda: list(DOMAINS)) shard_refresh_rows: int = 100_000 data_buf_size: int = 16_384 cpu_offload: bool = False use_compile: bool = True @property def dtype(self): return torch.bfloat16 @property def head_dim(self) -> int: return self.d_model // self.n_heads def effective_batch_tokens(self, ws: int) -> int: return self.micro_batch * self.grad_accum * ws * self.max_len def current_ctx_len(self, step: int) -> int: if self.ctx_warmup_steps <= 0 or step >= self.ctx_warmup_steps: return self.ctx_end_len t = step / self.ctx_warmup_steps raw = int(self.ctx_start_len + t * (self.ctx_end_len - self.ctx_start_len)) return ((raw + 63) // 64) * 64 # ══════════════════════════════════════════════════════════════════════════════ # FP8 HELPERS # FIX1: _has_fp8_support() checks major >= 9 so H100/H200 (sm_90) get fp8. # Old _is_blackwell() checked >= 10, silently killed fp8 on Hopper. # ══════════════════════════════════════════════════════════════════════════════ def _has_fp8_support() -> bool: """ FP8 tensor cores + torchao rowwise scaling are available on: Hopper sm_90 (H100, H200) <- primary target Blackwell sm_100+ (RTX PRO 6000) <- also supported Ampere sm_80 (A100) and below: NO fp8 tensor cores. """ if not torch.cuda.is_available(): return False dev = torch.cuda.current_device() major, _ = torch.cuda.get_device_capability(dev) return major >= 9 def resolve_fp8(cfg: NovaConfig) -> bool: if cfg.use_fp8 == "off": return False ok = HAS_TORCHAO and _has_fp8_support() if cfg.use_fp8 == "on" and not ok: log.warning( "[fp8] forced ON but torchao missing or GPU < sm_90 — bf16 fallback" ) return False if cfg.use_fp8 == "auto" and is_main(): if torch.cuda.is_available(): major, minor = torch.cuda.get_device_capability( torch.cuda.current_device() ) else: major, minor = 0, 0 log.info( f"[fp8] auto → torchao={HAS_TORCHAO} " f"compute_cap=sm_{major}{minor} " f"fp8_capable={_has_fp8_support()} " f"→ {'ON' if ok else 'OFF'}" ) return ok _FP8_SKIP = ( "embed", "lm_head", "norm", "diff_lambda", "head_gate", "ls_a", "ls_f", "rope", ) def _fp8_filter(module: nn.Module, fqn: str) -> bool: return ( isinstance(module, nn.Linear) and not any(s in fqn for s in _FP8_SKIP) ) def apply_fp8(model: nn.Module, cfg: NovaConfig) -> nn.Module: if not resolve_fp8(cfg): return model fp8_cfg = Float8LinearConfig.from_recipe_name("rowwise") convert_to_float8_training(model, config=fp8_cfg, module_filter_fn=_fp8_filter) if is_main(): n = sum( 1 for name, m in model.named_modules() if isinstance(m, nn.Linear) and _fp8_filter(m, name) ) log.info( f"[fp8] {n} Linear layers -> Float8Linear " f"(rowwise, Hopper/Blackwell)" ) return model # ══════════════════════════════════════════════════════════════════════════════ # DATA # ══════════════════════════════════════════════════════════════════════════════ class ShardStream: """ Per-domain infinite shard iterator. Periodically re-queries HF for new shards so training automatically picks up newly uploaded data files without restart — general and code shards trickling in are handled. """ def __init__(self, domain: str, refresh_every: int = 100_000): self.domain = domain self.refresh_every = refresh_every self._seen: List[str] = [] self._queue: List[str] = [] self._since_refresh = 0 self._refresh() def _refresh(self): all_shards = list_shards(self.domain) new = [s for s in all_shards if s not in self._seen] if new and is_main(): log.info(f"[data] {self.domain}: +{len(new)} new shards discovered") random.shuffle(new) self._queue.extend(new) self._seen.extend(new) def _load(self, path: str) -> np.ndarray: local = hf_download_dataset(path) arr = np.load(local, mmap_mode="r") idx = np.random.permutation(len(arr)) return arr[idx] def __iter__(self) -> Iterator[np.ndarray]: while True: if not self._queue: q = list(self._seen) random.shuffle(q) self._queue = q self._refresh() path = self._queue.pop(0) for row in self._load(path): yield row self._since_refresh += 1 if self._since_refresh >= self.refresh_every: self._since_refresh = 0 self._refresh() class BalancedIterableDataset(IterableDataset): def __init__( self, domains: List[str], seq_len: int, buf_size: int = 16_384, refresh_every: int = 100_000, ): super().__init__() self.domains = domains self.seq_len = seq_len self.buf_size = buf_size self.refresh_every = refresh_every def __iter__(self) -> Iterator[torch.Tensor]: streams = {d: iter(ShardStream(d, self.refresh_every)) for d in self.domains} per_domain = max(1, self.buf_size // len(self.domains)) while True: buf: List[np.ndarray] = [] for stream in streams.values(): for _ in range(per_domain): buf.append(next(stream)) random.shuffle(buf) for row in buf: tokens = row[: self.seq_len].astype(np.int64) if len(tokens) < self.seq_len: tokens = np.pad( tokens, (0, self.seq_len - len(tokens)), constant_values=0, ) yield torch.from_numpy(tokens.copy()) def make_loaders(cfg: NovaConfig, local_rank: int, ws: int): kw = dict( pin_memory=True, num_workers=4, persistent_workers=True, prefetch_factor=4, ) train_ds = BalancedIterableDataset( cfg.domains, cfg.max_len, cfg.data_buf_size, cfg.shard_refresh_rows ) val_ds = BalancedIterableDataset( cfg.domains, cfg.max_len, cfg.data_buf_size // 4, cfg.shard_refresh_rows ) train_loader = DataLoader(train_ds, batch_size=cfg.micro_batch, **kw) val_loader = DataLoader(val_ds, batch_size=min(cfg.micro_batch, 4), **kw) return train_loader, val_loader # ══════════════════════════════════════════════════════════════════════════════ # MODEL BUILDING BLOCKS # ══════════════════════════════════════════════════════════════════════════════ class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.scale = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: return F.rms_norm(x, self.scale.shape, self.scale, self.eps) def precompute_rope(head_dim: int, max_len: int, theta: float, device=None): inv_freq = 1.0 / ( theta ** ( torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim ) ) t = torch.arange(max_len, dtype=torch.float32, device=device) freqs = torch.outer(t, inv_freq) return torch.cos(freqs), torch.sin(freqs) def apply_rope(q, k, cos, sin): L = q.shape[1] cos_full = torch.cat([cos[:L], cos[:L]], dim=-1) sin_full = torch.cat([sin[:L], sin[:L]], dim=-1) cos_full = cos_full.unsqueeze(0).unsqueeze(2) sin_full = sin_full.unsqueeze(0).unsqueeze(2) def _rotate_half(x): half = x.shape[-1] // 2 x1, x2 = x[..., :half], x[..., half:] return torch.cat([-x2, x1], dim=-1) qf = q.float() kf = k.float() q_rot = (qf * cos_full + _rotate_half(qf) * sin_full).to(q.dtype) k_rot = (kf * cos_full + _rotate_half(kf) * sin_full).to(k.dtype) return q_rot, k_rot class NovaAttention(nn.Module): def __init__(self, cfg: NovaConfig): super().__init__() self.nh = cfg.n_heads self.nkv = cfg.n_kv_heads self.hd = cfg.head_dim self.ng = cfg.n_heads // cfg.n_kv_heads self.half = cfg.head_dim // 2 q_dim = cfg.n_heads * cfg.head_dim kv_dim = cfg.n_kv_heads * cfg.head_dim self.qkv = nn.Linear(cfg.d_model, q_dim + 2 * kv_dim, bias=False) self.o_proj = nn.Linear(q_dim, cfg.d_model, bias=False) self.q_norm = RMSNorm(cfg.head_dim) self.k_norm = RMSNorm(cfg.head_dim) self.h_norm = RMSNorm(cfg.head_dim) self.diff_lambda = nn.Parameter( torch.full((cfg.n_heads,), cfg.diff_lambda_init) ) self.head_gate = nn.Parameter(torch.ones(cfg.n_heads)) self.lambda_init = cfg.diff_lambda_init std = 0.02 dstd = std / math.sqrt(2 * cfg.n_layers) nn.init.normal_(self.qkv.weight, std=std) nn.init.normal_(self.o_proj.weight, std=dstd) def forward(self, x, cos, sin): B, L, D = x.shape qkv = self.qkv(x) q_dim = self.nh * self.hd kv_d = self.nkv * self.hd q, k, v = qkv.split([q_dim, kv_d, kv_d], dim=-1) q = q.view(B, L, self.nh, self.hd) k = k.view(B, L, self.nkv, self.hd) v = v.view(B, L, self.nkv, self.hd) q = self.q_norm(q) k = self.k_norm(k) q, k = apply_rope(q, k, cos, sin) q = q.transpose(1, 2) k = k.transpose(1, 2).repeat_interleave(self.ng, dim=1) v = v.transpose(1, 2).repeat_interleave(self.ng, dim=1) q1, q2 = q[..., :self.half], q[..., self.half:] k1, k2 = k[..., :self.half], k[..., self.half:] v1, v2 = v[..., :self.half], v[..., self.half:] out1 = F.scaled_dot_product_attention(q1, k1, v1, is_causal=True) out2 = F.scaled_dot_product_attention(q2, k2, v2, is_causal=True) lam = self.diff_lambda.view(1, self.nh, 1, 1).to(out1.dtype) out = torch.cat([out1, lam * out2], dim=-1) out = self.h_norm(out) * (1.0 - self.lambda_init) out = out * self.head_gate.view(1, self.nh, 1, 1).to(out.dtype) out = out.transpose(1, 2).contiguous().view(B, L, -1) return self.o_proj(out) class SwiGLU(nn.Module): def __init__(self, cfg: NovaConfig): super().__init__() h = cfg.ffn_hidden self.gate = nn.Linear(cfg.d_model, h, bias=False) self.up = nn.Linear(cfg.d_model, h, bias=False) self.down = nn.Linear(h, cfg.d_model, bias=False) std = 0.02 dstd = std / math.sqrt(2 * cfg.n_layers) nn.init.normal_(self.gate.weight, std=std) nn.init.normal_(self.up.weight, std=std) nn.init.normal_(self.down.weight, std=dstd) def forward(self, x): return self.down(F.silu(self.gate(x)) * self.up(x)) class NovaBlock(nn.Module): def __init__(self, cfg: NovaConfig): super().__init__() self.attn_norm = RMSNorm(cfg.d_model) self.ffn_norm = RMSNorm(cfg.d_model) self.attn = NovaAttention(cfg) self.ffn = SwiGLU(cfg) self.ls_a = nn.Parameter(torch.full((cfg.d_model,), 0.1)) self.ls_f = nn.Parameter(torch.full((cfg.d_model,), 0.1)) def forward(self, x, cos, sin): x = x + self.ls_a * self.attn(self.attn_norm(x), cos, sin) x = x + self.ls_f * self.ffn(self.ffn_norm(x)) return x class Nova10BDense(nn.Module): def __init__(self, cfg: NovaConfig): super().__init__() self.cfg = cfg self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model) self.layers = nn.ModuleList([NovaBlock(cfg) for _ in range(cfg.n_layers)]) self.norm = RMSNorm(cfg.d_model) self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) self.softcap = cfg.logit_softcap self._use_tied_lm_head = False nn.init.normal_(self.embed.weight, std=0.02) nn.init.normal_(self.lm_head.weight, std=0.02) cos, sin = precompute_rope(cfg.head_dim, cfg.max_len, cfg.rope_theta) self.register_buffer("rope_cos", cos, persistent=False) self.register_buffer("rope_sin", sin, persistent=False) if is_main(): n = self.count_params() log.info( f"[model] Nova-10B-Dense | {n / 1e9:.2f}B params | " f"{cfg.n_layers} layers | d={cfg.d_model} | " f"heads={cfg.n_heads}/{cfg.n_kv_heads} | " f"ffn_h={cfg.ffn_hidden} | max_len={cfg.max_len}" ) def forward(self, ids, targets=None): B, L = ids.shape x = self.embed(ids) cos = self.rope_cos[:L] sin = self.rope_sin[:L] for layer in self.layers: x = layer(x, cos, sin) x = self.norm(x) if self._use_tied_lm_head: logits = F.linear(x, self.embed.weight) else: logits = self.lm_head(x) logits = self.softcap * torch.tanh(logits / self.softcap) if targets is None: return logits loss = F.cross_entropy( logits[:, :-1].contiguous().view(-1, self.cfg.vocab_size), targets[:, 1:].contiguous().view(-1), ignore_index=-1, ) return logits, loss def count_params(self) -> int: seen, n = set(), 0 for p in self.parameters(): if id(p) not in seen: seen.add(id(p)) n += p.numel() return n def weight_checksum(self) -> float: """ Cheap deterministic checksum over a sample of parameters. Verifies all ranks built bit-identical models after seed_for_model_init without needing a full sync_module_states broadcast. """ with torch.no_grad(): s = 0.0 s += float(self.embed.weight[:100, :100].sum().item()) s += float(self.lm_head.weight[:100, :100].sum().item()) for i, layer in enumerate(self.layers): if i % 5 == 0: s += float(layer.attn.qkv.weight[:50, :50].sum().item()) s += float(layer.ffn.gate.weight[:50, :50].sum().item()) return s def build_model(cfg: NovaConfig) -> nn.Module: """ Build on CPU. Caller must have called seed_for_model_init(same_seed) on ALL ranks immediately before this so every rank's random init is bit-identical — no FSDP broadcast needed afterward. """ model = Nova10BDense(cfg).to(cfg.dtype) model = apply_fp8(model, cfg) return model def enable_weight_tying(fsdp_model: nn.Module): for module in fsdp_model.modules(): if isinstance(module, Nova10BDense): module._use_tied_lm_head = True if is_main(): log.info("[model] weight tying enabled (post-FSDP, flag-based)") return if is_main(): log.warning("[model] could not find Nova10BDense for weight tying") # ══════════════════════════════════════════════════════════════════════════════ # FSDP WRAPPING # ══════════════════════════════════════════════════════════════════════════════ _SAVE_SD_CFG = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) _LOAD_SD_CFG = FullStateDictConfig(offload_to_cpu=True, rank0_only=False) def wrap_fsdp(model: nn.Module, cfg: NovaConfig, device: torch.device) -> FSDP: mp_policy = MixedPrecision( param_dtype=torch.bfloat16, reduce_dtype=torch.bfloat16, buffer_dtype=torch.bfloat16, ) wrap_policy = functools.partial( transformer_auto_wrap_policy, transformer_layer_cls={NovaBlock}, ) fsdp_model = FSDP( model, auto_wrap_policy = wrap_policy, mixed_precision = mp_policy, sharding_strategy = ShardingStrategy.FULL_SHARD, backward_prefetch = BackwardPrefetch.BACKWARD_PRE, cpu_offload = CPUOffload(offload_params=cfg.cpu_offload), device_id = device, use_orig_params = True, # FIX3: sync_module_states=False — correctness guaranteed by identical # CPU seeding in seed_for_model_init(). Verified cheaply via # verify_model_sync() checksum all_reduce. The old True value caused # a giant NCCL broadcast at construction time that crashed on both # Blackwell and early Hopper runs. sync_module_states = False, ) if is_main(): log.info( f"[fsdp] FULL_SHARD × {world_size()} GPUs | " f"cpu_offload={cfg.cpu_offload} | bf16 mixed-precision | " f"sync_module_states=False (models pre-synced via seeding)" ) return fsdp_model # ══════════════════════════════════════════════════════════════════════════════ # MUON OPTIMIZER # ══════════════════════════════════════════════════════════════════════════════ def _newtonschulz5(G: torch.Tensor, steps: int = 5) -> torch.Tensor: assert G.ndim == 2, f"_newtonschulz5: expected 2-D, got shape {tuple(G.shape)}" a, b, c = 3.4445, -4.7750, 2.0315 X = G.bfloat16() if X.size(0) > X.size(1): X = X.T X = X / (X.norm() + 1e-7) for _ in range(steps): A = X @ X.T X = a * X + (b * A + c * (A @ A)) @ X if G.size(0) > G.size(1): X = X.T return X class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, momentum=0.95, weight_decay=0.0, ns_steps=5): defaults = dict( lr=lr, momentum=momentum, weight_decay=weight_decay, ns_steps=ns_steps ) super().__init__(list(params), defaults) @torch.no_grad() def step(self): for g in self.param_groups: lr = g["lr"] mom = g["momentum"] wd = g["weight_decay"] ns = g["ns_steps"] for p in g["params"]: if p.grad is None: continue grad = p.grad state = self.state[p] if grad.ndim != 2: if "buf" not in state: state["buf"] = torch.zeros_like(grad) buf = state["buf"] buf.mul_(mom).add_(grad) if wd > 0: p.mul_(1.0 - lr * wd) p.add_(buf, alpha=-lr) continue if "buf" not in state: state["buf"] = torch.zeros_like(grad) buf = state["buf"] buf.mul_(mom).add_(grad) grad_ns = grad.add(buf, alpha=mom) u = _newtonschulz5(grad_ns, steps=ns) if wd > 0: p.mul_(1.0 - lr * wd) scale = 0.2 * math.sqrt(max(p.size(0), p.size(1))) p.add_(u.to(p.dtype), alpha=-lr * scale) def build_optimizers(model: nn.Module, cfg: NovaConfig): """ FIX2: MUST be called on the raw unwrapped CPU model BEFORE wrap_fsdp(). After FSDP wraps the model, named_parameters() walks FSDP's internal flat-param structure. Parameter names look different and the 2D-weight classification logic can see an empty set, producing: ValueError: optimizer got an empty parameter list With use_orig_params=True FSDP keeps references to the original param objects alive, so optimizer references built pre-wrap remain valid and correct after wrapping. Always call: build_optimizers() -> wrap_fsdp(), never the other way around. """ muon_p = [] aw_decay = [] aw_nodecay = [] seen: set = set() for name, p in model.named_parameters(): if id(p) in seen or not p.requires_grad: continue seen.add(id(p)) is_embed_or_head = ("embed" in name) or ("lm_head" in name) if p.ndim == 2 and not is_embed_or_head: muon_p.append(p) elif p.ndim >= 2: aw_decay.append(p) else: aw_nodecay.append(p) if is_main(): log.info( f"[optim] param classification: " f"muon={len(muon_p)} tensors | " f"aw_decay={len(aw_decay)} | " f"aw_nodecay={len(aw_nodecay)}" ) # Hard guard — if this fires, call order is wrong if len(muon_p) == 0: raise RuntimeError( "[optim] muon_p is empty — no 2D non-embed/lm_head params found.\n" "build_optimizers() MUST be called on the unwrapped model BEFORE " "wrap_fsdp(). Check call order in train()." ) muon_opt = Muon( muon_p, lr=cfg.muon_lr, momentum=cfg.muon_momentum, weight_decay=cfg.weight_decay, ) aw_groups = [ {"params": aw_decay, "weight_decay": cfg.weight_decay}, {"params": aw_nodecay, "weight_decay": 0.0}, ] if HAS_BNB: aw_opt = bnb.optim.AdamW8bit( aw_groups, lr=cfg.lr, betas=(cfg.beta1, cfg.beta2), eps=1e-8, ) if is_main(): log.info("[optim] Muon + AdamW-8bit (bitsandbytes)") else: aw_opt = torch.optim.AdamW( aw_groups, lr=cfg.lr, betas=(cfg.beta1, cfg.beta2), eps=1e-8, fused=True, ) if is_main(): log.warning("[optim] bitsandbytes not found — fused fp32 AdamW") if is_main(): nm = sum(p.numel() for p in muon_p) na = sum(p.numel() for p in aw_decay + aw_nodecay) log.info( f"[optim] Muon {nm / 1e6:.0f}M params | " f"AdamW {na / 1e6:.0f}M params" ) return muon_opt, aw_opt def cosine_lr(step, warmup, total, peak, floor): if step < warmup: return peak * (step + 1) / warmup if step >= total: return floor t = (step - warmup) / (total - warmup) return floor + 0.5 * (peak - floor) * (1.0 + math.cos(math.pi * t)) # ══════════════════════════════════════════════════════════════════════════════ # CHECKPOINT I/O # ══════════════════════════════════════════════════════════════════════════════ def _full_state_dict_save(model: FSDP) -> dict: with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, _SAVE_SD_CFG): return model.state_dict() def save_ckpt(model, muon_opt, aw_opt, step, epoch, cfg, val_loss, tag) -> bool: barrier() sd = _full_state_dict_save(model) if not is_main(): return True base = f"{tag}_s{step:07d}" mpath = os.path.join(TEMP_CKPT_DIR, f"{base}.model.pt") xpath = os.path.join(TEMP_CKPT_DIR, f"{base}.meta.pt") try: torch.save({k: v.contiguous() for k, v in sd.items()}, mpath) torch.save( { "step": step, "epoch": epoch, "val_loss": val_loss, "config": asdict(cfg), "muon": muon_opt.state_dict(), "adamw": aw_opt.state_dict(), }, xpath, ) ok1 = hf_upload(mpath, f"{base}.model.pt") ok2 = hf_upload(xpath, f"{base}.meta.pt") for p in [mpath, xpath]: try: os.remove(p) except FileNotFoundError: pass if ok1 and ok2: rotate_ckpts(tag, keep=cfg.ckpt_keep) return ok1 and ok2 except Exception as e: log.error(f"[ckpt] save failed: {e}") return False def save_final(model, step, epoch, cfg, val_loss): sd = _full_state_dict_save(model) if not is_main(): return base = f"nova10b_dense_final_s{step:07d}" ext = "safetensors" if HAS_ST else "pt" mpath = os.path.join(TEMP_CKPT_DIR, f"{base}.{ext}") cpath = os.path.join(TEMP_CKPT_DIR, f"{base}_config.json") try: sd_clean = {k: v.detach().clone().contiguous() for k, v in sd.items()} if HAS_ST: st_save(sd_clean, mpath) else: torch.save(sd_clean, mpath) with open(cpath, "w") as f: json.dump( { **asdict(cfg), "final_step": step, "epoch": epoch, "val_loss": val_loss, }, f, indent=2, ) hf_upload(mpath, os.path.basename(mpath)) hf_upload(cpath, os.path.basename(cpath)) for p in [mpath, cpath]: try: os.remove(p) except FileNotFoundError: pass log.info(f"[final] {base} saved") except Exception as e: log.error(f"[final] save failed: {e}") def load_ckpt(model, muon_opt, aw_opt, device, cfg) -> tuple: try: files = hf_api.list_repo_files( repo_id=HF_MODEL_REPO, repo_type="model", token=HF_TOKEN, ) ckpts = sorted( [f for f in files if f.endswith(".model.pt")], key=_ckpt_step ) if not ckpts: return 0, 0, float("inf") latest = ckpts[-1] base = latest[: -len(".model.pt")] mpath = hf_download_ckpt(f"{base}.model.pt") xpath = hf_download_ckpt(f"{base}.meta.pt") if not mpath or not xpath: return 0, 0, float("inf") sd = torch.load(mpath, map_location="cpu", weights_only=True) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, _LOAD_SD_CFG): model.load_state_dict(sd, strict=False) barrier() meta = torch.load(xpath, map_location="cpu", weights_only=False) for opt, key in [(muon_opt, "muon"), (aw_opt, "adamw")]: if key in meta: try: opt.load_state_dict(meta[key]) except Exception as e: log.warning(f"[ckpt] {key} optimizer not restored: {e}") saved_cfg = meta.get("config", {}) for field_name in [ "total_steps", "warmup_steps", "ctx_warmup_steps", "ctx_start_len", "ctx_end_len", ]: if field_name in saved_cfg: setattr(cfg, field_name, saved_cfg[field_name]) s, e, v = meta["step"], meta["epoch"], meta["val_loss"] if is_main(): log.info(f"[ckpt] resumed step={s} epoch={e} val={v:.4f}") return s, e, v except Exception as ex: if is_main(): log.error(f"[ckpt] resume failed: {ex} — fresh start") return 0, 0, float("inf") # ══════════════════════════════════════════════════════════════════════════════ # VALIDATION # ══════════════════════════════════════════════════════════════════════════════ @torch.no_grad() def run_val(model, loader, device, n: int = 80) -> float: model.eval() losses: List[float] = [] it = iter(loader) for _ in range(n): try: b = next(it).to(device, non_blocking=True) except StopIteration: break with torch.amp.autocast("cuda", dtype=torch.bfloat16): _, loss = model(b, targets=b) losses.append(loss.item()) model.train() val = float(np.mean(losses)) if losses else float("nan") local = torch.tensor(val, device=device) all_reduce_mean(local) return local.item() # ══════════════════════════════════════════════════════════════════════════════ # UTILS # ══════════════════════════════════════════════════════════════════════════════ def vram_str(dev: Optional[int] = None) -> str: if dev is None: dev = torch.cuda.current_device() used = torch.cuda.memory_allocated(dev) / 1024 ** 3 total = torch.cuda.get_device_properties(dev).total_memory / 1024 ** 3 return f"{used:.1f}/{total:.0f} GB" # ══════════════════════════════════════════════════════════════════════════════ # TRAIN # # Stage order (IMPORTANT — do not reorder without reading FIX2 comment): # seed_model_init # config_build # vocab_metadata # model_build_cpu # verify_model_sync # optimizer_build <- BEFORE fsdp_wrap (FIX2) # fsdp_wrap <- AFTER optimizer_build (FIX2) # compile # seed_runtime # checkpoint_resume # data_loaders # first_step_smoke_test # ══════════════════════════════════════════════════════════════════════════════ def train(): local_rank = setup_dist() device = torch.device(f"cuda:{local_rank}") ws = world_size() log_cuda_environment(local_rank) log_p2p_matrix(local_rank, ws) with _stage("cuda_health_check"): cuda_health_check(device) with _stage("nccl_collective_selftest"): nccl_collective_selftest(device, ws) with _stage("seed_model_init"): seed_for_model_init(42) init_hf() barrier() with _stage("config_build"): cfg = NovaConfig() target_tokens = 200_000_000_000 tokens_per_step = cfg.effective_batch_tokens(ws) cfg.total_steps = max( cfg.total_steps, math.ceil(target_tokens / tokens_per_step), ) if is_main(): log.info( f"[schedule] {ws} GPUs | " f"{tokens_per_step / 1e6:.1f}M tok/step | " f"{cfg.total_steps} steps ≈ " f"{cfg.total_steps * tokens_per_step / 1e9:.0f}B tokens" ) with _stage("vocab_metadata"): if is_main(): try: meta_p = hf_download_dataset("metadata.json") with open(meta_p) as f: dmeta = json.load(f) actual = dmeta.get("vocab_size", cfg.vocab_size) cfg.vocab_size = (actual + 63) // 64 * 64 log.info(f"[data] vocab {actual} -> {cfg.vocab_size}") except Exception as e: log.info( f"[data] no metadata.json ({e}) " f"— default vocab={cfg.vocab_size}" ) vs = torch.tensor(cfg.vocab_size, device=device) dist.broadcast(vs, src=0) cfg.vocab_size = int(vs.item()) with _stage("model_build_cpu"): if is_main(): log.info("[model] building Nova-10B-Dense on CPU …") with torch.device("cpu"): model = build_model(cfg) with _stage("verify_model_sync"): # Cheap correctness check replacing sync_module_states broadcast. # Sampled param checksum all_reduce — if all ranks seeded identically # every checksum is bit-identical. local_checksum = ( model.weight_checksum() if hasattr(model, "weight_checksum") else 0.0 ) checksum_t = torch.tensor( local_checksum, dtype=torch.float64, device=device ) gathered = [torch.zeros_like(checksum_t) for _ in range(ws)] dist.all_gather(gathered, checksum_t) vals = [float(g.item()) for g in gathered] max_diff = max(vals) - min(vals) if is_main(): if max_diff < 1e-6: log.info( f"[verify] model checksums MATCH across all {ws} ranks " f"(checksum={vals[0]:.6f}) — no broadcast needed ✓" ) else: log.error( f"[verify] model checksums DIVERGE across ranks: {vals} " f"(max_diff={max_diff:.6f}) — ranks built DIFFERENT models! " f"Investigate seed_for_model_init before continuing." ) # ── FIX2: optimizer built HERE on raw CPU model, BEFORE fsdp_wrap ──────── # After wrap_fsdp(), named_parameters() walks FSDP's flat-param structure # and the 2D-weight classifier finds nothing -> empty muon_p -> ValueError. # With use_orig_params=True the optimizer's param references stay valid # after wrapping, so this ordering is both correct and safe. with _stage("optimizer_build"): muon_opt, aw_opt = build_optimizers(model, cfg) with _stage("fsdp_wrap"): model = wrap_fsdp(model, cfg, device) enable_weight_tying(model) with _stage("compile"): if cfg.use_compile: try: for m in model.modules(): if isinstance(m, Nova10BDense): m.forward = torch.compile( m.forward, mode="max-autotune", fullgraph=False, dynamic=True, ) if is_main(): log.info( "[compile] Nova10BDense.forward compiled " "(max-autotune, dynamic)" ) break except Exception as e: if is_main(): log.warning(f"[compile] failed, running eager: {e}") if is_main(): log.info(f"[model] VRAM after wrap: {vram_str()}") with _stage("seed_runtime"): seed_for_runtime(42 + rank(), device) with _stage("checkpoint_resume"): start_step, start_epoch, best_val = 0, 0, float("inf") try: start_step, start_epoch, best_val = load_ckpt( model, muon_opt, aw_opt, device, cfg ) except Exception as e: if is_main(): log.error(f"[ckpt] resume error: {e} — fresh start") info = torch.tensor( [start_step, start_epoch], dtype=torch.long, device=device ) dist.broadcast(info, src=0) start_step, start_epoch = int(info[0]), int(info[1]) with _stage("data_loaders"): train_loader, val_loader = make_loaders(cfg, rank(), ws) train_iter = itertools.cycle(train_loader) with _stage("first_step_smoke_test"): try: probe_batch = next(train_iter).to(device, non_blocking=True) ctx0 = cfg.current_ctx_len(0) if ctx0 < probe_batch.shape[1]: probe_batch = probe_batch[:, :ctx0] with torch.amp.autocast("cuda", dtype=cfg.dtype): _, probe_loss = model(probe_batch, targets=probe_batch) (probe_loss / cfg.grad_accum).backward() torch.cuda.synchronize() if is_main(): log.info( f"[smoke] first forward+backward OK, " f"loss={probe_loss.item():.4f}" ) muon_opt.zero_grad(set_to_none=True) aw_opt.zero_grad(set_to_none=True) except Exception as e: log.error( f"[smoke] rank={rank()} FIRST forward/backward failed: {e}" ) raise if is_main(): log.info("=" * 72) log.info( f"Nova-10B-Dense | H200 × {ws} | " f"FP8={resolve_fp8(cfg)} | FSDP FULL_SHARD | " f"Muon + {'AdamW-8bit' if HAS_BNB else 'AdamW-fp32'}" ) log.info( f" d={cfg.d_model} layers={cfg.n_layers} " f"heads={cfg.n_heads}/{cfg.n_kv_heads} " f"ffn_h={cfg.ffn_hidden} max_len={cfg.max_len}" ) log.info( f" micro={cfg.micro_batch} × accum={cfg.grad_accum} × " f"{ws} GPUs × seq={cfg.max_len} = " f"{tokens_per_step / 1e6:.1f}M tok/step" ) log.info( f" {cfg.total_steps} steps -> " f"≈{cfg.total_steps * tokens_per_step / 1e9:.0f}B tokens" ) log.info( f" muon_lr={cfg.muon_lr} aw_lr={cfg.lr} " f"warmup={cfg.warmup_steps} budget={cfg.max_hours}h" ) log.info(f" resume step={start_step} epoch={start_epoch}") log.info("=" * 72) model.train() muon_opt.zero_grad(set_to_none=True) aw_opt.zero_grad(set_to_none=True) step = start_step epoch = start_epoch t_start = time.time() t_last_ckpt = time.time() losses: List[float] = [] accum_loss = 0.0 accum_n = 0 def elapsed_h() -> float: return (time.time() - t_start) / 3600 def remaining_h() -> float: return cfg.max_hours - elapsed_h() while step < cfg.total_steps: if elapsed_h() >= cfg.max_hours: if is_main(): log.info(f"[time] {cfg.max_hours}h budget exhausted — stopping") break try: batch = next(train_iter) except Exception as e: if is_main(): log.error(f"[data] fetch error: {e}") continue batch = batch.to(device, non_blocking=True) ctx_raw = cfg.current_ctx_len(step) ctx_t = torch.tensor(ctx_raw, device=device) dist.broadcast(ctx_t, src=0) ctx = int(ctx_t.item()) if ctx < batch.shape[1]: batch = batch[:, :ctx] try: with torch.amp.autocast("cuda", dtype=cfg.dtype): _, loss = model(batch, targets=batch) (loss / cfg.grad_accum).backward() accum_loss += loss.detach().item() accum_n += 1 except torch.cuda.OutOfMemoryError: torch.cuda.empty_cache() muon_opt.zero_grad(set_to_none=True) aw_opt.zero_grad(set_to_none=True) accum_loss = accum_n = 0 if is_main(): log.error( "[oom] step skipped — reduce micro_batch. " "(Unlikely on 141GB H200 HBM3e unless micro_batch is huge.)" ) continue except RuntimeError as e: log.error( f"[train_step] rank={rank()} device={device} step={step} " f"CUDA RuntimeError: {e}" ) raise if accum_n < cfg.grad_accum: continue mlr = cosine_lr( step, cfg.warmup_steps, cfg.total_steps, cfg.muon_lr, cfg.muon_min_lr ) alr = cosine_lr( step, cfg.warmup_steps, cfg.total_steps, cfg.lr, cfg.min_lr ) for pg in muon_opt.param_groups: pg["lr"] = mlr for pg in aw_opt.param_groups: pg["lr"] = alr torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip) muon_opt.step() aw_opt.step() muon_opt.zero_grad(set_to_none=True) aw_opt.zero_grad(set_to_none=True) avg = accum_loss / cfg.grad_accum losses.append(avg) if len(losses) > 500: losses.pop(0) accum_loss = accum_n = 0 step += 1 if is_main() and step % cfg.log_every == 0: sps = (step - start_step) / max(time.time() - t_start, 1e-6) tps = sps * tokens_per_step a50 = float(np.mean(losses[-50:])) if losses else 0.0 ppl = math.exp(min(avg, 20)) log.info(f"step {step:06d} | ctx {ctx} | {remaining_h():.2f}h left") log.info( f" loss {avg:.4f} avg50 {a50:.4f} ppl {ppl:.1f}" f" mlr {mlr:.2e} alr {alr:.2e}" ) log.info( f" {tps / 1e3:6.0f}K tok/s " f"{1 / max(sps, 1e-9):.2f}s/step " f"vram {vram_str()}" ) if step % cfg.eval_every == 0: vl = run_val(model, val_loader, device, n=80) vpp = math.exp(min(vl, 20)) if is_main(): log.info(f" >> val {vl:.4f} ppl {vpp:.2f}") if vl < best_val: best_val = vl save_ckpt( model, muon_opt, aw_opt, step, epoch, cfg, vl, tag="best", ) log.info(" >> new best checkpoint saved") if step % cfg.ckpt_every == 0: save_ckpt( model, muon_opt, aw_opt, step, epoch, cfg, losses[-1] if losses else 0.0, tag="ckpt", ) if time.time() - t_last_ckpt > 1800: save_ckpt( model, muon_opt, aw_opt, step, epoch, cfg, losses[-1] if losses else 0.0, tag="latest", ) t_last_ckpt = time.time() if is_main(): log.info("=" * 72) log.info("TRAINING COMPLETE — saving final checkpoint …") fval = run_val(model, val_loader, device, n=200) barrier() save_final(model, step, epoch, cfg, fval) if is_main(): avg_t = float(np.mean(losses[-200:])) if losses else 0.0 report = dict( step = step, epoch = epoch, train_loss = avg_t, val_loss = fval, val_ppl = math.exp(min(fval, 20)), best_val = best_val, elapsed_h = elapsed_h(), tokens_seen = step * tokens_per_step, gpus = ws, fp8 = resolve_fp8(cfg), compile = cfg.use_compile, optimizer = "Muon+AdamW8bit" if HAS_BNB else "Muon+AdamW", model = "Nova-10B-Dense", hardware = f"{ws}x H200 141GB HBM3e", ) rpath = os.path.join(TEMP_CKPT_DIR, f"report_s{step}.json") with open(rpath, "w") as f: json.dump(report, f, indent=2) hf_upload(rpath, f"report_s{step}.json") try: os.remove(rpath) except FileNotFoundError: pass log.info("=" * 72) for k, v in report.items(): log.info( f" {k:<22}: " + (f"{v:.4f}" if isinstance(v, float) else str(v)) ) log.info("=" * 72) cleanup_dist() if __name__ == "__main__": train()