| |
| """ |
| 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 |
| """ |
|
|
| |
| |
| |
| 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") |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
|
|
| 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) |
|
|
|
|
| |
| |
| |
|
|
| 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 |
| 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_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}") |
|
|
|
|
| |
| |
| |
|
|
| @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 |
|
|
|
|
| |
| |
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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") |
|
|
|
|
| |
| |
| |
|
|
| _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, |
| |
| |
| |
| |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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)}" |
| ) |
|
|
| |
| 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)) |
|
|
|
|
| |
| |
| |
|
|
| 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") |
|
|
|
|
| |
| |
| |
|
|
| @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() |
|
|
|
|
| |
| |
| |
|
|
| 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" |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| 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"): |
| |
| |
| |
| 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." |
| ) |
|
|
| |
| |
| |
| |
| |
| 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() |