From 2821bd646faada79508313e835456500a37dbd3d Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Mon, 10 Aug 2026 16:31:50 +0800 Subject: [PATCH] late loss.to("CPU") in llama (#17476) * late loss.to("CPU") in llama * acc = 0 --- examples/mlperf/model_train.py | 21 +++++++++++---------- 1 file changed, 11 insertions(+), 10 deletions(-) diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index e9d8086aab..47c4977e39 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -1458,7 +1458,8 @@ def train_llama3(): # realize everything here if optim.master_params: Tensor.realize(*optim.master_params) - Tensor.realize(*optim.params, *fp8_inv_scales, *fp8_amax, *fp8_next_amax, *fp8_grad_amax, *fp8_next_grad_amax) + loss_acc = Tensor.zeros(1, dtype=dtypes.float32, device=device) + Tensor.realize(loss_acc, *optim.params, *fp8_inv_scales, *fp8_amax, *fp8_next_amax, *fp8_grad_amax, *fp8_next_grad_amax) @TinyJit def minibatch(tokens:Tensor): @@ -1476,8 +1477,8 @@ def train_llama3(): 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, *fp8_amax, *fp8_next_amax, *fp8_grad_amax, *fp8_next_grad_amax) + loss_acc.assign(loss_acc + loss.flatten().float()) + return loss_acc.realize(*grads, *fp8_amax, *fp8_next_amax, *fp8_grad_amax, *fp8_next_grad_amax) @TinyJit def optim_step(): @@ -1490,9 +1491,10 @@ def train_llama3(): 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, *fp8_amax, *fp8_grad_amax) + loss_cpu = loss_acc.to("CPU") + Tensor.realize(lr_cpu, grad_norm_cpu, loss_cpu, loss_acc.assign(0), *grads, *fp8_inv_scales, *fp8_amax, *fp8_grad_amax) - return lr_cpu, grad_norm_cpu + return lr_cpu, grad_norm_cpu, loss_cpu @TinyJit @Context(TRAINING=0) @@ -1547,8 +1549,8 @@ def train_llama3(): st = time.perf_counter() stopped = False - losses, data_time, dev_time = [], 0, 0 - for _ in range(grad_acc if i >= 2 else 1): + data_time, dev_time = 0, 0 + for _ in range(accum_steps:=grad_acc if i >= 2 else 1): ist = time.perf_counter() try: tokens = next(train_iter) except StopIteration: @@ -1556,16 +1558,15 @@ def train_llama3(): break mst = time.perf_counter() data_time += mst - ist - losses.append(minibatch(tokens).item()) + minibatch(tokens) 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() + lr, grad_norm, loss = ret[0].item(), ret[1].item(), ret[2].item() / accum_steps et = time.perf_counter() - loss = sum(losses) / len(losses) optim_time = et - gt dev_time += optim_time step_time = et - st