#!/usr/bin/env python3 # ruff: noqa: E501 """Stage 2: single-node/multi-node NVIDIA GPU port of the Marin/Levanter Qwen3 SFT recipe. PyTorch + HF transformers + torch FSDP. No JAX, no GCS, no Iris, no marin credentials. Reads the packed dataset produced by prepare_sft_data.py, writes native (resumable) checkpoints and HF-format exports to local disk on the same cadence as the TPU run, so the existing evalchemy tooling can consume `/hf/step-N` unchanged. Launch (8 GPUs, one node): torchrun --standalone --nproc_per_node=8 train_sft_qwen3.py --data /data/sft/foo --out /data/runs/foo Plan only, no GPUs needed: python train_sft_qwen3.py --data /data/sft/foo --out /tmp/x --dry-run WHAT IS REPRODUCED EXACTLY -------------------------- * optimiser chain order: global-norm clip -> Adam(b1,b2,eps) -> decoupled wd -> scale(-lr) * weight-decay mask: no decay on *norm*, embeddings and biases; lm_head IS decayed * LR: linear warmup 0->lr over int(warmup*T) steps, constant, then cosine to lr*min_lr_ratio over int(decay*T) steps (clamped). LR at optimiser step 0 is exactly 0.0. * mixed precision: fp32 master params / bf16 compute / fp32 gradient reduction * loss: per-microbatch token-weighted mean over assistant tokens, then arithmetic mean across the `--loss-groups` microbatches of an optimiser step (levanter grad_accum uses ReductionType.MEAN over 2 microbatches of 32 on v5p-64). No z-loss. * RoPE positions are CONTIGUOUS 0..L-1 across a packed example (levanter passes `pos_ids = arange(Pos)`; it does NOT reset per document), while attention is still blocked across documents with varlen cu_seq_lens. Most HF/TRL packing implementations reset position_ids per document -- that is a real divergence and we do not do it. * data order: sequential over packed examples with wraparound (shuffle is a silent no-op in the TPU config), so `--data-seed` only matters if you pass --shuffle. * checkpoint/export naming: levanter's StepInfo.step == completed_steps - 1, and hooks fire when `step > 1 and step % every == 0`, plus a forced final hook. A 2000-step run therefore writes hf/step-100 ... hf/step-1900 and a final hf/step-1999. WHAT DIFFERS (and why) -- see also the handoff doc: * per-device microbatch is 1 packed example (not 32-across-32-chips); grouping into `--loss-groups` groups restores the exact loss normalisation, but the *gradient* accumulation order differs, so results are not bitwise identical. * attention kernel is FlashAttention-2 (varlen), not the TPU splash/Pallas kernel. * cross-entropy is a chunked torch implementation, not levanter's fused Pallas CE. * HF export defaults to bf16 here (the TPU export was fp32; vLLM casts to bf16 anyway). """ from __future__ import annotations import argparse import contextlib import json import math import os import shutil import time from dataclasses import dataclass, replace import numpy as np MODEL_VOCAB_SIZE = 151936 # Qwen3 config.vocab_size for 1.7B / 4B / 8B / 32B # ---------------------------------------------------------------------------------- # recipes (frozen per model size, copied from the marin experiment files) # ---------------------------------------------------------------------------------- @dataclass class Recipe: model_id: str learning_rate: float weight_decay: float warmup: float decay: float min_lr_ratio: float max_grad_norm: float beta1: float beta2: float eps: float train_batch_size: int num_train_steps: int max_seq_len: int steps_per_checkpoint: int steps_per_hf_export: int lr_schedule: str = "cosine" RECIPES: dict[str, Recipe] = { # experiments/exp_sft_qwen3_8b_selfinstill_*_2k_lr5e6_wd01.py "qwen3-8b": Recipe( model_id="Qwen/Qwen3-8B", learning_rate=5e-6, weight_decay=0.01, warmup=0.05, decay=0.9, min_lr_ratio=0.1, max_grad_norm=1.0, beta1=0.9, beta2=0.999, eps=1e-8, train_batch_size=64, num_train_steps=2000, max_seq_len=32768, steps_per_checkpoint=20, steps_per_hf_export=100, ), # experiments/exp_*_sft_qwen3_4b_* (OT4 recipe: higher LR, no wd, tight clip) "qwen3-4b": Recipe( model_id="Qwen/Qwen3-4B", learning_rate=2e-5, weight_decay=0.0, warmup=0.03, decay=0.9, min_lr_ratio=0.1, max_grad_norm=0.2, beta1=0.9, beta2=0.999, eps=1e-8, train_batch_size=128, num_train_steps=0, # 0 => ceil(8 * n_packs / batch); override with --num-train-steps max_seq_len=32768, steps_per_checkpoint=20, steps_per_hf_export=100, ), "qwen3-1.7b": Recipe( model_id="Qwen/Qwen3-1.7B", learning_rate=2e-5, weight_decay=0.0, warmup=0.03, decay=0.9, min_lr_ratio=0.1, max_grad_norm=0.2, beta1=0.9, beta2=0.999, eps=1e-8, train_batch_size=128, num_train_steps=0, max_seq_len=32768, steps_per_checkpoint=20, steps_per_hf_export=100, ), } # ---------------------------------------------------------------------------------- # LR schedule (port of levanter OptimizerConfig.lr_scheduler for a single cycle) # ---------------------------------------------------------------------------------- def _frac_or_steps(v: float, total: int) -> int: if v < 0.0 or (v > 1.0 and v % 1 != 0): raise ValueError(f"Invalid fraction {v}") return int(v * total) if v <= 1.0 else int(v) def make_lr_multiplier(recipe: Recipe, num_train_steps: int): warmup_steps = min(_frac_or_steps(recipe.warmup, num_train_steps), num_train_steps) max_decay = max(num_train_steps - warmup_steps, 0) decay_steps = min(max(_frac_or_steps(recipe.decay, num_train_steps), 0), max_decay) stable_steps = num_train_steps - warmup_steps - decay_steps alpha = recipe.min_lr_ratio sched = recipe.lr_schedule def fn(step: int) -> float: if warmup_steps > 0 and step < warmup_steps: return step / warmup_steps # optax.linear_schedule(0.0, lr, warmup_steps) if step < warmup_steps + stable_steps: return 1.0 if decay_steps == 0: return 1.0 t = min(step - warmup_steps - stable_steps, decay_steps) if sched == "cosine": cos = 0.5 * (1.0 + math.cos(math.pi * t / decay_steps)) return (1.0 - alpha) * cos + alpha if sched == "linear": return 1.0 + (alpha - 1.0) * (t / decay_steps) if sched == "constant": return 1.0 raise ValueError(f"unsupported lr_schedule {sched}") return fn, (warmup_steps, stable_steps, decay_steps) # ---------------------------------------------------------------------------------- # weight decay mask (port of AdamConfig.build_weight_decay_mask reasonable_default) # ---------------------------------------------------------------------------------- def is_no_decay(name: str) -> bool: """levanter excludes LayerNorm/RMSNorm/RmsNorm/Embedding modules and *.bias. lm_head is an ordinary Linear in levanter, so it IS decayed. With tied embeddings (1.7B/4B) the single shared tensor is named model.embed_tokens.weight and is excluded, which matches levanter (lm_head is None when tie_word_embeddings=True). """ if name.endswith(".bias"): return True if "norm" in name.lower(): # input_layernorm, post_attention_layernorm, model.norm, q_norm, k_norm return True if name.endswith("embed_tokens.weight"): return True return False # ---------------------------------------------------------------------------------- # packed dataset # ---------------------------------------------------------------------------------- class PackedSFTData: def __init__(self, path: str): with open(os.path.join(path, "manifest.json")) as f: self.manifest = json.load(f) self.tokens = np.load(os.path.join(path, "tokens.npy"), mmap_mode="r") self.mask = np.load(os.path.join(path, "assistant_mask.npy"), mmap_mode="r") self.doc_lens = np.load(os.path.join(path, "doc_lens.npy")) self.pack_offsets = np.load(os.path.join(path, "pack_offsets.npy")) self.assistant_counts = np.load(os.path.join(path, "assistant_counts.npy")) self.max_seq_len = int(self.manifest["max_seq_len"]) self.n_packs = int(self.tokens.shape[0]) assert self.tokens.shape[1] == self.max_seq_len def doc_lengths(self, i: int) -> list[int]: s, e = int(self.pack_offsets[i]), int(self.pack_offsets[i + 1]) return [int(x) for x in self.doc_lens[s:e]] def collate(data: PackedSFTData, indices, device): """Flatten `indices` packed examples into one padding-free row of B*L tokens. Returns tensors shaped [1, B*L] plus varlen metadata. batch dim is always 1 so the HF flash-attention path uses flash_attn_varlen_func with our explicit cu_seq_lens. """ import torch L = data.max_seq_len ids = np.concatenate([np.asarray(data.tokens[i], dtype=np.int64) for i in indices]) msk = np.concatenate([np.asarray(data.mask[i], dtype=np.int64) for i in indices]) labels = np.where(msk == 1, ids, -100) for b in range(len(indices)): labels[b * L] = -100 # levanter's not_last_mask: position L-1 of the PREVIOUS example targets = np.empty_like(labels) targets[:-1] = labels[1:] targets[-1] = -100 seg_lens: list[int] = [] for i in indices: dl = data.doc_lengths(i) seg_lens.extend(dl) pad = L - int(sum(dl)) if pad > 0: seg_lens.append(pad) # padding is its own attention segment cu = np.concatenate([[0], np.cumsum(np.asarray(seg_lens, dtype=np.int64))]).astype(np.int32) pos = np.tile(np.arange(L, dtype=np.int64), len(indices)) # contiguous per example, like levanter n_targets = int((targets != -100).sum()) return { "input_ids": torch.from_numpy(ids).unsqueeze(0).to(device, non_blocking=True), "position_ids": torch.from_numpy(pos).unsqueeze(0).to(device, non_blocking=True), "targets": torch.from_numpy(targets).unsqueeze(0).to(device, non_blocking=True), "cu_seq_lens": torch.from_numpy(cu).to(device, non_blocking=True), "max_length": int(max(seg_lens)), }, n_targets # ---------------------------------------------------------------------------------- # model wrapper with chunked (memory-bounded) cross entropy # ---------------------------------------------------------------------------------- def _build_sft_module(): import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint def _ce_chunk(hidden_chunk, target_chunk, lm_head): logits = lm_head(hidden_chunk).float() return F.cross_entropy(logits, target_chunk, ignore_index=-100, reduction="sum") class SFTModule(nn.Module): """hf_model + a chunked lm_head/CE that never materialises [L, 151936] logits.""" def __init__(self, hf_model, loss_chunk_size: int, checkpoint_loss: bool = True): super().__init__() self.hf_model = hf_model self.loss_chunk_size = loss_chunk_size self.checkpoint_loss = checkpoint_loss def forward(self, input_ids, position_ids, targets, cu_seq_lens, max_length, return_hidden: bool = False): out = self.hf_model.model( input_ids=input_ids, position_ids=position_ids, use_cache=False, cu_seq_lens_q=cu_seq_lens, cu_seq_lens_k=cu_seq_lens, max_length_q=max_length, max_length_k=max_length, ) hidden = out.last_hidden_state[0] if return_hidden: return hidden tgt = targets[0] total = torch.zeros((), device=hidden.device, dtype=torch.float32) n = hidden.shape[0] cs = self.loss_chunk_size or n for s in range(0, n, cs): e = min(s + cs, n) if self.checkpoint_loss and self.training: part = checkpoint(_ce_chunk, hidden[s:e], tgt[s:e], self.hf_model.lm_head, use_reentrant=False) else: part = _ce_chunk(hidden[s:e], tgt[s:e], self.hf_model.lm_head) total = total + part return total return SFTModule # ---------------------------------------------------------------------------------- # distributed helpers # ---------------------------------------------------------------------------------- def dist_info(): return ( int(os.environ.get("RANK", 0)), int(os.environ.get("LOCAL_RANK", 0)), int(os.environ.get("WORLD_SIZE", 1)), ) def log0(rank: int, msg: str) -> None: if rank == 0: print(msg, flush=True) def fmt_hms(seconds: float) -> str: seconds = max(int(seconds), 0) h, rem = divmod(seconds, 3600) m, s = divmod(rem, 60) return f"{h}:{m:02d}:{s:02d}" def detect_grad_reduce_factor(device, world_size: int, rank: int) -> float: """FSDP averages gradients over the DP group; loss scales must undo that. Empirically determined with a 1-layer probe so a torch behaviour change cannot silently rescale the effective learning rate by world_size. """ import torch import torch.nn as nn from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy if world_size == 1: return 1.0 lin = nn.Linear(64, 64, bias=False) with torch.no_grad(): lin.weight.zero_() probe = FSDP( lin, sharding_strategy=ShardingStrategy.FULL_SHARD, device_id=device, use_orig_params=True, ) x = torch.ones(1, 64, device=device) probe(x).sum().backward() with FSDP.summon_full_params(probe, with_grads=True, writeback=False): observed = float(lin.weight.grad.detach().float().mean().item()) del probe, lin if abs(observed - 1.0) < 1e-3: log0(rank, "[init] FSDP averages gradients across ranks -> loss scale multiplied by world_size") return float(world_size) if abs(observed - world_size) < 1e-3: log0(rank, "[init] FSDP sums gradients across ranks -> no loss rescale") return 1.0 raise RuntimeError(f"Unexpected FSDP gradient reduction: mean grad {observed}, world_size {world_size}") def verify_document_isolation(model, device, rank: int) -> None: """Prove cross-document attention is actually blocked before burning GPU-days. If transformers stops forwarding cu_seq_lens_* to the attention kernel this test fails loudly instead of silently training with cross-document leakage. """ import torch L, d0 = 64, 32 ids = torch.randint(0, 1000, (1, L), device=device) pos = torch.arange(L, device=device).unsqueeze(0) cu = torch.tensor([0, d0, L], dtype=torch.int32, device=device) tgt = torch.full((1, L), -100, dtype=torch.long, device=device) was_training = model.training model.eval() with torch.no_grad(): h1 = model(ids, pos, tgt, cu, d0, return_hidden=True)[:d0].clone() ids2 = ids.clone() ids2[0, d0:] = torch.randint(0, 1000, (L - d0,), device=device) h2 = model(ids2, pos, tgt, cu, d0, return_hidden=True)[:d0].clone() if was_training: model.train() diff = (h1 - h2).abs().max().item() if diff != 0.0: raise RuntimeError( "Cross-document attention is NOT blocked (max hidden-state delta " f"{diff:.3e} after perturbing the second document). The varlen cu_seq_lens " "kwargs are not reaching flash-attention. Check the transformers version " "and that attn_implementation='flash_attention_2'." ) log0(rank, "[init] cross-document attention isolation verified (delta == 0)") # ---------------------------------------------------------------------------------- def build_model(args, recipe: Recipe, rank: int, local_rank: int, world_size: int): import torch from torch.distributed.fsdp import BackwardPrefetch, MixedPrecision, ShardingStrategy from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp.wrap import ModuleWrapPolicy from transformers import AutoConfig, AutoModelForCausalLM config = AutoConfig.from_pretrained(recipe.model_id) config.use_cache = False if config.vocab_size != MODEL_VOCAB_SIZE: log0(rank, f"[init] WARNING model vocab_size={config.vocab_size}, expected {MODEL_VOCAB_SIZE}") if args.init_mode == "rank0_meta" and world_size > 1: if rank == 0: hf_model = AutoModelForCausalLM.from_pretrained( recipe.model_id, dtype=torch.float32, attn_implementation=args.attn_impl, low_cpu_mem_usage=True ) else: with torch.device("meta"): hf_model = AutoModelForCausalLM.from_config(config, attn_implementation=args.attn_impl) else: hf_model = AutoModelForCausalLM.from_pretrained( recipe.model_id, dtype=torch.float32, attn_implementation=args.attn_impl, low_cpu_mem_usage=True ) hf_model.config.use_cache = False hf_model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) sft = _build_sft_module()(hf_model, args.loss_chunk_size, checkpoint_loss=not args.no_loss_checkpoint) layer_cls = type(hf_model.model.layers[0]) wrap_classes = {layer_cls} if args.wrap_embeddings: wrap_classes.add(torch.nn.Embedding) mp = MixedPrecision( param_dtype=torch.bfloat16, reduce_dtype=torch.float32 if args.grad_reduce_dtype == "float32" else torch.bfloat16, buffer_dtype=torch.float32, ) def param_init_fn(module): module.to_empty(device=torch.device("cuda", local_rank), recurse=False) model = FSDP( sft, auto_wrap_policy=ModuleWrapPolicy(wrap_classes), mixed_precision=mp, sharding_strategy=ShardingStrategy.FULL_SHARD, backward_prefetch=BackwardPrefetch.BACKWARD_PRE, device_id=torch.device("cuda", local_rank), use_orig_params=True, limit_all_gathers=True, sync_module_states=(args.init_mode == "rank0_meta" and world_size > 1), param_init_fn=param_init_fn if (args.init_mode == "rank0_meta" and world_size > 1 and rank != 0) else None, ) return model, hf_model def build_optimizer(model, recipe: Recipe, args): import torch decay, no_decay = [], [] dn, nn_ = [], [] for name, p in model.named_parameters(): if not p.requires_grad: continue clean = name.replace("_fsdp_wrapped_module.", "").replace("hf_model.", "") (no_decay if is_no_decay(clean) else decay).append(p) (nn_ if is_no_decay(clean) else dn).append(clean) groups = [ {"params": decay, "weight_decay": recipe.weight_decay}, {"params": no_decay, "weight_decay": 0.0}, ] opt = torch.optim.AdamW( groups, lr=recipe.learning_rate, betas=(recipe.beta1, recipe.beta2), eps=recipe.eps, fused=args.fused_adam, foreach=not args.fused_adam, ) return opt, dn, nn_ # ---------------------------------------------------------------------------------- # checkpoint / export # ---------------------------------------------------------------------------------- def save_native(model, optimizer, step: int, path: str, extra: dict) -> None: import torch.distributed as dist import torch.distributed.checkpoint as dcp from torch.distributed.checkpoint.state_dict import get_state_dict msd, osd = get_state_dict(model, optimizer) tmp = path + ".partial" dcp.save({"model": msd, "optim": osd}, checkpoint_id=tmp) dist.barrier() if dist.get_rank() == 0: with open(os.path.join(tmp, "meta.json"), "w") as f: json.dump({"step": step, **extra}, f) if os.path.exists(path): shutil.rmtree(path) os.replace(tmp, path) dist.barrier() def load_native(model, optimizer, path: str) -> int: import torch.distributed.checkpoint as dcp from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict msd, osd = get_state_dict(model, optimizer) state = {"model": msd, "optim": osd} dcp.load(state, checkpoint_id=path) set_state_dict(model, optimizer, model_state_dict=state["model"], optim_state_dict=state["optim"]) with open(os.path.join(path, "meta.json")) as f: return int(json.load(f)["step"]) def export_hf(model, hf_model, tokenizer, out_dir: str, label: int, dtype_name: str, keep_gen_cfg: bool) -> None: import torch import torch.distributed as dist from torch.distributed.fsdp import FullStateDictConfig, StateDictType from torch.distributed.fsdp import FullyShardedDataParallel as FSDP cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, cfg): sd = model.state_dict() if dist.get_rank() == 0: dtype = {"bfloat16": torch.bfloat16, "float32": torch.float32, "float16": torch.float16}[dtype_name] clean = {} for k, v in sd.items(): k = k[len("hf_model.") :] if k.startswith("hf_model.") else k clean[k] = v.to(dtype) target = os.path.join(out_dir, "hf", f"step-{label}") tmp = target + ".partial" if os.path.exists(tmp): shutil.rmtree(tmp) os.makedirs(tmp, exist_ok=True) hf_model.save_pretrained(tmp, state_dict=clean, safe_serialization=True, max_shard_size="5GB") tokenizer.save_pretrained(tmp) cfg_path = os.path.join(tmp, "config.json") with open(cfg_path) as f: cfg_json = json.load(f) cfg_json["torch_dtype"] = dtype_name # vLLM reads this; keep it honest cfg_json["dtype"] = dtype_name with open(cfg_path, "w") as f: json.dump(cfg_json, f, indent=2) gen_cfg = os.path.join(tmp, "generation_config.json") if not keep_gen_cfg and os.path.exists(gen_cfg): os.remove(gen_cfg) # the levanter export writes no generation_config.json if os.path.exists(target): shutil.rmtree(target) os.replace(tmp, target) print(f"[export] wrote {target}", flush=True) del sd dist.barrier() def export_tokenizer(recipe: Recipe): """267 <|padding_i|> tokens so len(tokenizer) == model vocab, exactly like HFCheckpointConverter.with_tokenizer_padded_to_match_model().""" from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(recipe.model_id) missing = MODEL_VOCAB_SIZE - len(tok) if missing > 0: tok.add_tokens([f"<|padding_{i}|>" for i in range(missing)]) return tok, missing # ---------------------------------------------------------------------------------- def dry_run(args, recipe: Recipe, data: PackedSFTData, num_train_steps: int) -> None: lr_fn, (w, s, d) = make_lr_multiplier(recipe, num_train_steps) L = recipe.max_seq_len tokens_per_step = recipe.train_batch_size * L print("=" * 78) print(f"model {recipe.model_id}") print(f"packed examples {data.n_packs} (seq len {L})") print(f"global batch {recipe.train_batch_size} examples = {tokens_per_step:,} tokens/step") print(f"steps {num_train_steps} => {tokens_per_step * num_train_steps:,} tokens presented") print(f"epochs over the data {num_train_steps * recipe.train_batch_size / data.n_packs:.2f}") print(f"lr schedule warmup {w} / stable {s} / cosine {d} to {recipe.learning_rate * recipe.min_lr_ratio:.3e}") for st in [0, 1, w - 1, w, w + s - 1, w + s, num_train_steps // 2, num_train_steps - 1]: if 0 <= st < num_train_steps: print(f" step {st:>5} lr = {recipe.learning_rate * lr_fn(st):.6e}") print(f"weight decay {recipe.weight_decay} (excluded: *norm*, embed_tokens, *.bias; lm_head decayed)") print(f"grad clip / betas {recipe.max_grad_norm} / ({recipe.beta1}, {recipe.beta2}) eps {recipe.eps}") print(f"loss groups per step {args.loss_groups} (TPU used 2 microbatches of 32)") print(f"assistant tokens {int(data.assistant_counts.sum()):,} over the dataset") print(f"exports hf/step-N every {recipe.steps_per_hf_export}, native every {recipe.steps_per_checkpoint}") print("=" * 78) # ---------------------------------------------------------------------------------- def main() -> None: p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) p.add_argument("--data", required=True, help="output dir of prepare_sft_data.py") p.add_argument("--out", required=True, help="run dir; checkpoints/, checkpoints-temp/, hf/ live here") p.add_argument("--recipe", default="qwen3-8b", choices=sorted(RECIPES)) p.add_argument("--model-id", default=None, help="override the recipe's HF model id") p.add_argument("--num-train-steps", type=int, default=0) p.add_argument("--train-batch-size", type=int, default=0) p.add_argument("--per-device-batch", type=int, default=1, help="packed examples per GPU per micro-step") p.add_argument("--loss-groups", type=int, default=2, help="microbatches per optimiser step; TPU used 2") p.add_argument("--steps-per-checkpoint", type=int, default=0) p.add_argument("--steps-per-hf-export", type=int, default=0) p.add_argument("--temp-checkpoint-minutes", type=float, default=10.0) p.add_argument("--loss-chunk-size", type=int, default=2048) p.add_argument("--no-loss-checkpoint", action="store_true") p.add_argument("--attn-impl", default="flash_attention_2") p.add_argument("--grad-reduce-dtype", default="float32", choices=("float32", "bfloat16")) p.add_argument("--init-mode", default="rank0_meta", choices=("rank0_meta", "all_ranks")) p.add_argument("--wrap-embeddings", action="store_true") p.add_argument("--fused-adam", action="store_true") p.add_argument("--hf-export-dtype", default="bfloat16", choices=("bfloat16", "float32", "float16")) p.add_argument("--keep-generation-config", action="store_true") p.add_argument("--trainer-seed", type=int, default=0, help="levanter TrainerConfig.seed (stays 0 on TPU)") p.add_argument("--data-seed", type=int, default=42, help="only used when --shuffle is set") p.add_argument("--shuffle", action="store_true", help="NOT what the TPU run did; shuffle is a no-op there") p.add_argument("--resume", default="auto", help="auto | none | ") p.add_argument("--log-every", type=int, default=1) p.add_argument("--skip-checks", action="store_true") p.add_argument("--wandb-project", default=None) p.add_argument("--wandb-name", default=None) p.add_argument("--wandb-id", default=None, help="fixed W&B run id; reused across restarts (with resume=allow) so an " "auto-resumed run continues ONE W&B run instead of forking a new one") p.add_argument("--dry-run", action="store_true") args = p.parse_args() recipe = RECIPES[args.recipe] if args.model_id: recipe = replace(recipe, model_id=args.model_id) if args.train_batch_size: recipe = replace(recipe, train_batch_size=args.train_batch_size) if args.steps_per_checkpoint: recipe = replace(recipe, steps_per_checkpoint=args.steps_per_checkpoint) if args.steps_per_hf_export: recipe = replace(recipe, steps_per_hf_export=args.steps_per_hf_export) data = PackedSFTData(args.data) recipe = replace(recipe, max_seq_len=data.max_seq_len) num_train_steps = args.num_train_steps or recipe.num_train_steps if not num_train_steps: num_train_steps = math.ceil(8 * data.n_packs / recipe.train_batch_size) if args.dry_run: dry_run(args, recipe, data, num_train_steps) return import torch import torch.distributed as dist rank, local_rank, world_size = dist_info() torch.cuda.set_device(local_rank) dist.init_process_group("nccl") device = torch.device("cuda", local_rank) torch.manual_seed(args.trainer_seed) np.random.seed(args.trainer_seed) B = recipe.train_batch_size G = args.loss_groups per_group = B // G micro_per_group = per_group // (world_size * args.per_device_batch) if B % G or per_group % (world_size * args.per_device_batch): raise ValueError( f"train_batch_size={B} must divide by loss_groups={G} and then by " f"world_size*per_device_batch={world_size * args.per_device_batch}" ) if recipe.steps_per_hf_export <= 0 or recipe.steps_per_checkpoint <= 0: raise ValueError("steps_per_hf_export and steps_per_checkpoint must be > 0") log0(rank, f"[init] world_size={world_size} batch={B} groups={G} micro/group={micro_per_group} seq={recipe.max_seq_len}") grad_factor = 1.0 if args.skip_checks else detect_grad_reduce_factor(device, world_size, rank) model, hf_model = build_model(args, recipe, rank, local_rank, world_size) if not args.skip_checks: verify_document_isolation(model, device, rank) optimizer, decayed, undecayed = build_optimizer(model, recipe, args) log0(rank, f"[init] weight decay applied to {len(decayed)} tensors, excluded {len(undecayed)}") log0(rank, f"[init] excluded sample: {undecayed[:4]} ... lm_head decayed: {any(n.endswith('lm_head.weight') for n in decayed)}") lr_fn, (w_steps, s_steps, d_steps) = make_lr_multiplier(recipe, num_train_steps) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_fn) tokenizer, n_pad = export_tokenizer(recipe) log0(rank, f"[init] export tokenizer padded with {n_pad} tokens -> len={len(tokenizer)}") ckpt_dir = os.path.join(args.out, "checkpoints") temp_dir = os.path.join(args.out, "checkpoints-temp") if rank == 0: os.makedirs(ckpt_dir, exist_ok=True) os.makedirs(os.path.join(args.out, "hf"), exist_ok=True) dist.barrier() start_step = 0 resume_path = None if args.resume == "auto": cands = [] if os.path.exists(os.path.join(temp_dir, "meta.json")): cands.append(temp_dir) if os.path.isdir(ckpt_dir): cands += [os.path.join(ckpt_dir, d) for d in os.listdir(ckpt_dir) if d.startswith("step-")] cands = [c for c in cands if os.path.exists(os.path.join(c, "meta.json"))] if cands: resume_path = max(cands, key=lambda c: json.load(open(os.path.join(c, "meta.json")))["step"]) elif args.resume != "none": resume_path = args.resume if resume_path: start_step = load_native(model, optimizer, resume_path) for _ in range(start_step): scheduler.step() log0(rank, f"[init] resumed from {resume_path} at step {start_step}") order = np.arange(data.n_packs) if args.shuffle: np.random.default_rng(args.data_seed).shuffle(order) run = None if args.wandb_project and rank == 0: import wandb run = wandb.init( project=args.wandb_project, name=args.wandb_name or os.path.basename(args.out), id=args.wandb_id, resume="allow" if args.wandb_id else None, ) run.config.update({**vars(args), **recipe.__dict__, "num_train_steps": num_train_steps}, allow_val_change=True) model.train() t_start = time.time() last_temp = time.time() log0(rank, "First batch loaded, starting first train step (includes CUDA graph/kernel warmup)...") for step in range(start_step, num_train_steps): t_step = time.time() base = step * B batch_idx = [int(order[(base + i) % data.n_packs]) for i in range(B)] optimizer.zero_grad(set_to_none=True) group_loss_sums = torch.zeros(G, device=device, dtype=torch.float32) group_tokens = torch.zeros(G, device=device, dtype=torch.float32) for g in range(G): gidx = batch_idx[g * per_group : (g + 1) * per_group] denom = float(sum(int(data.assistant_counts[i]) for i in gidx)) if denom <= 0: raise RuntimeError(f"group {g} of step {step} has zero assistant tokens") scale = grad_factor / (G * denom) for m in range(micro_per_group): off = (m * world_size + rank) * args.per_device_batch mine = gidx[off : off + args.per_device_batch] batch, n_tok = collate(data, mine, device) loss_sum = model( batch["input_ids"], batch["position_ids"], batch["targets"], batch["cu_seq_lens"], batch["max_length"], ) (loss_sum * scale).backward() group_loss_sums[g] += loss_sum.detach() group_tokens[g] += n_tok dist.all_reduce(group_loss_sums) dist.all_reduce(group_tokens) if step == start_step: expect = [float(sum(int(data.assistant_counts[i]) for i in batch_idx[g * per_group : (g + 1) * per_group])) for g in range(G)] got = [float(x) for x in group_tokens.tolist()] if any(abs(a - b) > 0.5 for a, b in zip(expect, got)): raise RuntimeError(f"loss denominator mismatch: manifest={expect} observed={got}") log0(rank, f"[init] loss denominators verified: {got}") loss_value = float((group_loss_sums / group_tokens.clamp(min=1)).mean().item()) gnorm = model.clip_grad_norm_(recipe.max_grad_norm) optimizer.step() scheduler.step() completed = step + 1 label = completed - 1 # levanter StepInfo.step dt = time.time() - t_step if rank == 0 and (completed % args.log_every == 0 or completed == num_train_steps): done = completed - start_step total = num_train_steps - start_step rate = (time.time() - t_start) / max(done, 1) print( f"Progress on:train {completed}it/{num_train_steps / 1000:.2f}kit " f"rate:{rate:.1f}s/it remaining:{fmt_hms(rate * (total - done))} " f"elapsed:{fmt_hms(time.time() - t_start)} postfix:loss={loss_value:.3f}", flush=True, ) if run is not None: run.log( { "train/loss": loss_value, "train/lr": scheduler.get_last_lr()[0], "train/grad_norm": float(gnorm), "train/tokens": int(group_tokens.sum().item()), "train/step_time": dt, }, step=completed, ) is_final = completed == num_train_steps if (label > 1 and label % int(recipe.steps_per_hf_export) == 0) or is_final: export_hf(model, hf_model, tokenizer, args.out, label, args.hf_export_dtype, args.keep_generation_config) if (label > 1 and label % int(recipe.steps_per_checkpoint) == 0) or is_final: save_native(model, optimizer, completed, os.path.join(ckpt_dir, f"step-{label}"), {"label": label}) elif (time.time() - last_temp) > args.temp_checkpoint_minutes * 60: save_native(model, optimizer, completed, temp_dir, {"label": label}) last_temp = time.time() log0(rank, f"[done] {num_train_steps} steps in {fmt_hms(time.time() - t_start)}") if run is not None: run.finish() dist.barrier() dist.destroy_process_group() if __name__ == "__main__": with contextlib.suppress(KeyboardInterrupt): main()