model.train()
while step < args.max_steps:
    while step < args.max_steps: # 读取一个批次（batch）的数据
        for k in list(batch.keys()):            
            batch[k] = batch[k].to(device, non_blocking=True) # # 将数据送入模型
            # ……
            loss = outputs.loss / args.grad_accum # 计算损失（loss）
            
        scaler.scale(loss).backward() # 反向传播更新模型参数
        if is_main and step % args.log_every == 0:
            # 每 log_every 步打印当前损失和学习率
            elapsed = time.time() - start
            print(f"step {step} | loss {running_loss/args.log_every:.4f} | lr {scheduler.get_last_lr()[0]:.2e} | elapsed {elapsed:.1f}s")
        if step % args.save_every == 0:
            # 每 save_every 步保存一个检查点，包括模型权重、优化器、调度器、scaler状态
            save_ckpt(model.module if isinstance(model, nn.parallel.DistributedDataParallel) else model,
              optimizer, scheduler, scaler, step, args.output_dir, is_main)
