From a805ce03b12729fcdd6219ff7f00923fc533d508 Mon Sep 17 00:00:00 2001 From: wozeparrot Date: Tue, 7 Jul 2026 01:07:50 -0400 Subject: [PATCH] gptoss: model train (#16884) --- examples/mlperf/model_train.py | 270 +++++++++++++++++++++++++++++++++ 1 file changed, 270 insertions(+) diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index bf68e4fff1..78c331883c 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -1658,6 +1658,276 @@ def train_llama3(): if MLLOGGER and RUNMLPERF: MLLOGGER.start(key=mllog_constants.BLOCK_START, metadata={mllog_constants.SAMPLES_COUNT: sequences_seen}) +def train_gptoss(): + from examples.mlperf.models.gpt_oss import GPTOSS, GPT_OSS_20B, apply_grad, FP8_DTYPE + from examples.mlperf.lr_schedulers import CosineAnnealingLRWithWarmup + from examples.mlperf.optim import GradAccClipAdamW + + BENCHMARK = getenv("BENCHMARK") + + config = {} + BASEDIR = config["BASEDIR"] = Path(getenv("BASEDIR", "/raid/datasets/c4-8b/")) + BS = config["BS"] = getenv("BS", 16) + grad_acc = config["GRADIENT_ACC_STEPS"] = getenv("GRADIENT_ACC_STEPS", 1) + GBS = config["GLOBAL_BATCH_SIZE"] = BS * grad_acc + SEED = config["SEED"] = getenv("SEED", 5760) + DATA_SEED = config["DATA_SEED"] = getenv("DATA_SEED", SEED) + SEQLEN = config["SEQLEN"] = getenv("SEQLEN", 8192) + TRAIN_ON_VAL = config["TRAIN_ON_VAL"] = getenv("TRAIN_ON_VAL", 0) + MAX_STEPS = config["MAX_STEPS"] = getenv("MAX_STEPS", 1_200_000) + SAMPLES = config["SAMPLES"] = getenv("SAMPLES", 5_760 if TRAIN_ON_VAL else MAX_STEPS * GBS) + EVAL_SAMPLES = config["EVAL_SAMPLES"] = getenv("EVAL_SAMPLES", 1024) + WARMUP_STEPS = config["WARMUP_STEPS"] = getenv("WARMUP_STEPS", 128) + LR = config["LR"] = getenv("LR", 4e-4 * GBS / 16) + END_LR = config["END_LR"] = getenv("END_LR", 4e-5) + EVAL_FREQ = config["EVAL_FREQ"] = getenv("EVAL_FREQ", 12288) + EVAL_BS = config["EVAL_BS"] = getenv("EVAL_BS", 16) + EVAL_TARGET = config["EVAL_TARGET"] = getenv("EVAL_TARGET", 3.34) + + opt_adamw_beta_1 = 0.9 + opt_adamw_beta_2 = 0.95 + opt_adamw_epsilon = 1e-5 + opt_adamw_weight_decay = 0.1 + + opt_learning_rate_warmup_steps = WARMUP_STEPS + opt_learning_rate_decay_steps = MAX_STEPS - opt_learning_rate_warmup_steps + opt_base_learning_rate = LR + opt_end_learning_rate = END_LR + + Tensor.manual_seed(SEED) # seed for weight initialization + + # ** init wandb ** + WANDB = getenv("WANDB") + if WANDB: + import wandb + wandb_args = {"id": wandb_id, "resume": "must"} if (wandb_id := getenv("WANDB_RESUME", "")) else {} + wandb.init(config=config, **wandb_args, project="MLPerf-gpt-oss") + + model_params = GPT_OSS_20B + model_params['vocab_size'] = 128256 + real_vocab_size = model_params['vocab_size'] + if (layers:=getenv("LAYERS")) != 0: model_params['n_layers'] = layers + print(f"model parameters: {model_params}") + + model = GPTOSS(**model_params, max_context=SEQLEN) + + params = get_parameters(model) + + if getenv("EMPTYWEIGHT"): + for v in get_parameters(model): + v = v.assign(Tensor.empty(v.shape, dtype=v.dtype)) + + is_dp = (DP := getenv("DP", 1)) > 1 + is_sharding = is_dp + device_count = DP + device = tuple(f"{Device.DEFAULT}:{i}" for i in range(device_count)) + + model.shard(device, False) + + is_offload_optim = bool(getenv("OFFLOAD_OPTIM")) + is_fake_offload = Device.DEFAULT == "NULL" + optim_device = ("CPU" if not is_fake_offload else "NULL:99") if is_offload_optim else None + optim = GradAccClipAdamW(params, lr=0.0, b1=opt_adamw_beta_1, b2=opt_adamw_beta_2, eps=opt_adamw_epsilon, weight_decay=opt_adamw_weight_decay, grad_acc=grad_acc, device=optim_device) + + for p in optim.params: + grad_dtype = dtypes.bfloat16 if p.dtype == FP8_DTYPE else p.dtype + p.grad = p.zeros_like(dtype=grad_dtype).contiguous() + grads = [p.grad for p in optim.params] + + from extra.gemm.cdna_asm_gemm import _mx_block_scale + model_state = get_state_dict(model) + fp8_scale_names = {n: f"{n}_scale" for n, t in model_state.items() if t.dtype == FP8_DTYPE} + fp8_inv_scales = [model_state[sname] for sname in fp8_scale_names.values()] + for wname, sname in fp8_scale_names.items(): + w, scale = model_state[wname], model_state[sname] + w._inv_scale = scale + if optim.master_params: + master = optim.master_params[next(j for j, p in enumerate(optim.params) if p is w)] + inv = scale if scale.device == master.device else scale.to(master.device) + bs = _mx_block_scale(inv.reshape(-1, inv.shape[-1])).reshape(w.shape) + master.assign((master * bs).contiguous()) + + scheduler = CosineAnnealingLRWithWarmup(optim, opt_base_learning_rate, opt_end_learning_rate, opt_learning_rate_warmup_steps, opt_learning_rate_decay_steps) + + # realize everything here + if optim.master_params: Tensor.realize(*optim.master_params) + Tensor.realize(*optim.params, *fp8_inv_scales) + + @TinyJit + def minibatch(tokens:Tensor): + if is_dp: tokens = tokens.to(None).shard(device, 0) + if not is_sharding: tokens = tokens.to(None) + logits:Tensor = model(tokens[:, :-1], save=True) + loss = logits.sparse_categorical_crossentropy(tokens[:, 1:]) + + for g, new_g in zip(grads, loss.gradient(*optim.params)): + apply_grad(g, new_g.uop) + + loss_cpu = loss.flatten().float().to("CPU") + return loss_cpu.realize(*grads) + + @TinyJit + def optim_step(): + grad_norm = optim.fstep(grads) + scheduler.step() + + for g in grads: g.assign(0) + + lr_cpu = optim.lr.float().to("CPU") + grad_norm_cpu = grad_norm.float().to("CPU") + Tensor.realize(lr_cpu, grad_norm_cpu, *grads, *fp8_inv_scales) + + return lr_cpu, grad_norm_cpu + + @TinyJit + @Context(TRAINING=0) + def eval_step(tokens:Tensor): + if is_dp: tokens = tokens.to(None).shard(device, 0) + if not is_sharding: tokens = tokens.to(None) + logits:Tensor = model(tokens[:, :-1]) + loss = logits.sparse_categorical_crossentropy(tokens[:, 1:]) + return loss.flatten().float().to("CPU") + + # ** data iters ** + def fake_data(bs, samples): + import numpy as np + for _ in range(samples // bs): + fake_data_np = np.random.randint(0, real_vocab_size, size=(bs, SEQLEN + 1), dtype=np.int32) + yield Tensor(fake_data_np, device="NPY") + + def get_train_iter(): + if getenv("FAKEDATA", 0): + return fake_data(BS, SAMPLES) + else: + from examples.mlperf.dataloader import batch_load_llama3 + return batch_load_llama3(BS, SAMPLES, SEQLEN, BASEDIR, seed=DATA_SEED, val=bool(TRAIN_ON_VAL), small=True) + + if getenv("FAKEDATA", 0): + eval_dataset = None + else: + from examples.mlperf.dataloader import get_llama3_dataset + eval_dataset = get_llama3_dataset(EVAL_SAMPLES, SEQLEN, BASEDIR, val=True, small=True) + + def get_eval_iter(): + if eval_dataset is None: + return fake_data(EVAL_BS, EVAL_SAMPLES) + from examples.mlperf.dataloader import iterate_llama3_dataset + return iterate_llama3_dataset(eval_dataset, EVAL_BS) + + num_params = sum(p.numel() for p in params) - model_params["vocab_size"]*model_params["dim"] + train_iter = get_train_iter() + i, sequences_seen = 0, 0 + step_times = [] + + while i < MAX_STEPS: + GlobalCounters.reset() + actual_gbs = GBS if i >= 2 else BS + if getenv("TRAIN", 1): + profile_marker(f"train @ {i}") + st = time.perf_counter() + + stopped = False + losses, data_time, dev_time = [], 0, 0 + for _ in range(grad_acc if i >= 2 else 1): + ist = time.perf_counter() + try: tokens = next(train_iter) + except StopIteration: + stopped = True + break + mst = time.perf_counter() + data_time += mst - ist + losses.append(minibatch(tokens).item()) + dev_time += time.perf_counter() - mst + if stopped: break + + gt = time.perf_counter() + ret = optim_step() + lr, grad_norm = ret[0].item(), ret[1].item() + et = time.perf_counter() + + loss = sum(losses) / len(losses) + optim_time = et - gt + dev_time += optim_time + step_time = et - st + gbs_time = gt - st + if BENCHMARK: step_times.append(step_time) + + i += 1 + sequences_seen += actual_gbs + + mem_gb = GlobalCounters.mem_used / 1e9 + gflops = GlobalCounters.global_ops / 1e9 / dev_time + mfu = ((6 * num_params * SEQLEN * GBS) / (dev_time * device_count * 4.6e15)) * 100 + tqdm.write( + f"{i:5} {step_time:.3f} s step, {gbs_time:.3f} s gbs, {optim_time:.3f} s optim, {data_time:.3f} s data, {loss:.4f} loss, " \ + f"{lr:.12f} LR, {grad_norm:.6f} grad_norm, {mem_gb:.2f} GB used, {gflops:9.2f} GFLOPS, {mfu:5.2f}% MFU") + if DEBUG >= 1: tqdm.write(" mem per device: " + ', '.join(f"{dev}: {mem/1e9:.2f} GB" for dev, mem in sorted(GlobalCounters.mem_used_per_device.items()))) + + if WANDB: + wandb.log({ + "train/loss": loss, + "train/lr": lr, + "train/grad_norm": grad_norm, + "train/step_time": step_time, + "train/gbs_time": gbs_time, + "train/optim_time": optim_time, + "train/dev_time": dev_time, + "train/data_time": data_time, + "train/mem": mem_gb, + "train/GFLOPS": gflops, + "train/MFU": mfu, + "train/sequences_seen": sequences_seen + }) + + if (ckpt_freq := getenv("CKPT")) and (i % ckpt_freq == 0 and (i != 1 or ckpt_freq == 1)): + tqdm.write("saving checkpoint") + if not os.path.exists(ckpt_dir := "./ckpts"): os.mkdir(ckpt_dir) + fn = f"{ckpt_dir}/gptoss_{i}.safe" + safe_save(get_state_dict(model), fn) + + tqdm.write("saving optim checkpoint") + fn = f"{ckpt_dir}/gptoss_{i}_optim.safe" + safe_save(get_state_dict(scheduler), fn) + + if i == BENCHMARK: + median_step_time = sorted(step_times)[BENCHMARK // 2] + estimated_steps = MAX_STEPS + estimated_total_minutes = int(median_step_time * estimated_steps / 60) + print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") + print(f"epoch global_ops: {GlobalCounters.global_ops:_}, " + f"epoch global_mem: {GlobalCounters.global_mem:_}") + + if (sequences_seen // EVAL_FREQ != (sequences_seen - actual_gbs) // EVAL_FREQ and (i != 1 or EVAL_FREQ == 1)) or (BENCHMARK and i == BENCHMARK): + if EVAL_BS == 0: return + tqdm.write(f"evaluating after {sequences_seen} sequences") + profile_marker(f"eval @ {i}") + + # run eval + eval_losses = [] + eval_iter = get_eval_iter() + tqdm.write(f"evaluating {EVAL_SAMPLES//EVAL_BS} batches of {EVAL_BS} sequences") + + for j,tokens in tqdm(enumerate(eval_iter), total=EVAL_SAMPLES//EVAL_BS): + eval_losses += eval_step(tokens).tolist() + + if BENCHMARK and (j+1) == min(BENCHMARK, EVAL_SAMPLES//EVAL_BS): + return + + log_perplexity = sum(eval_losses) / len(eval_losses) + + tqdm.write(f"eval log perplexity: {log_perplexity:.4f}") + + if WANDB: + wandb.log({"eval/log_perplexity": log_perplexity, "eval/sequences_seen": sequences_seen}) + + if log_perplexity < EVAL_TARGET: + tqdm.write(f"target achieved after {sequences_seen} sequences") + if getenv("CKPT"): + if not os.path.exists(ckpt_dir := "./ckpts"): os.mkdir(ckpt_dir) + fn = f"{ckpt_dir}/gptoss.safe" + safe_save(get_state_dict(model), fn) + break + def train_stable_diffusion(): from extra.models.unet import UNetModel from examples.mlperf.dataloader import batch_load_train_stable_diffusion