TRAINING_TIME_SECONDS = 300.0 while runtime_seconds < TRAINING_TIME_SECONDS: step_started = time.perf_counter() input_batch, target_batch = next(train_iterator) optimizer.zero_grad(set_to_none=True) logits = model(input_batch) loss = F.cross_entropy( logits.flatten(0, 1), target_batch.flatten(), ) loss.backward() optimizer.step() runtime_seconds += time.perf_counter() - step_started optimizer_steps += 1 tokens_seen += input_batch.numel() # Periodic evaluation and logging omitted. __ __