nova-10b-dense / nova10b_dense.py
Bc-AI's picture
Update nova10b_dense.py
aeec779 verified
Raw
History Blame Contribute Delete
65.8 kB
#!/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()