svd-code / gpu-sft /scripts /gpu_sft /train_sft_qwen3.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
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)
# ----------------------------------------------------------------------------------
@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 | <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()