def train_step(model, device, optimizer, batch): # copy data to device batch = copy_to_device(batch, device) zero_grads(model) with torch.amp.autocast('cuda', dtype=torch.bfloat16): outputs = model(**batch) loss = outputs.loss loss.backward() mark_static_address(optimizer) optimizer.step() return loss