# full_response, response_mask, old_log_probs, advantages <--- Buffer # Recompute the new log_probs. Notice no torch.no_grad(), so gradients WILL BE USED here. logits = llm(input_ids=full_response).logits # Extract log probs from the logits # Does log_softmax over the vocabulary and extracts the log-prob of each selected token log_probs = calculate_log_probs( logits, full_responses ) # Calculate the clipped surrogate loss reasoning_loss = calculate_ppo_loss( log_probs, # Trainable old_log_probs, # Obtained from exploration, not trainable advantages, # Obtained from environment, not trainable response_mask # Obtained from exploration, not trainable ) # Optimizaiton steps accelerator.backward(reasoning_loss) optimizer.step() optimizer.zero_grad()