def train_one_context(train_context, device=DEVICE): batch_size = BATCH_SIZE_BY_CONTEXT[train_context] # Build loaders and restore the shared model initialization. model = GPTModel(MODEL_CONFIG) model.load_state_dict(SHARED_INITIAL_STATE) model = model.to(device) for step in range(NUM_TRAINING_STEPS): 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, loss, throughput, and memory tracking omitted. return result, history, generation_rows __ __