Download gpu-sft/scripts/gpu_sft/train_sft_qwen3.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 36.1 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/gpu-sft/scripts/gpu_sft/train_sft_qwen3.py
- Command line
-
hf download hf://fzzhang/svd-code/gpu-sft/scripts/gpu_sft/train_sft_qwen3.py
-
curl -L -o train_sft_qwen3.py https://huggingface.co/fzzhang/svd-code/resolve/main/gpu-sft/scripts/gpu_sft/train_sft_qwen3.py
36.1 kB
| #!/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 `<out>/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) | |
| # ---------------------------------------------------------------------------------- | |
| 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 | <path>") | |
| 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() | |