def mark_dynamic(batch): for key in ['decoder_input_ids', 'labels', 'decoder_attention_mask']: torch._dynamo.mark_dynamic(batch[key], 1) def train_step(model, device, optimizer, batch): # copy data to device batch = copy_to_device(batch, device) # mark inputs as dynamic to avoid recompilation mark_dynamic(batch) optimizer.zero_grad() with torch.amp.autocast('cuda', dtype=torch.bfloat16): outputs = model(**batch) loss = outputs.loss loss.backward() optimizer.step() return loss