Build A Reasoning Model From Scratch 6: Reinforcement Learning 1 (Implementing GRPO for RLVR)
Summary
This video provides a deep technical dive into training a reasoning model using Reinforcement Learning with Verifiable Rewards (RLVR), specifically implementing the Group Relative Policy Optimization (GRPO) algorithm. The process focuses on optimizing the model weights based solely on the final, correct answer (accuracy reward), rather than intermediate reasoning steps. Key steps include generating diverse model rollouts, calculating relative rewards and advantages, computing sequence log probabilities, and finally, implementing the GRPO loss function within a standard PyTorch training loop. The speaker emphasizes the computational cost and resource requirements of this process.
Key takeaways
-
RLVR vs. RLHF
30:38
RLVR (Reinforcement Learning with Verifiable Rewards) is simpler and more direct than RLHF (Reinforcement Learning with Human Feedback) because it replaces the expensive human reward model with a deterministic verifier (e.g., a math verifier) that grades the final answer's correctness.
-
GRPO Algorithm
38:24
GRPO is presented as a simpler and less memory-intensive alternative to standard PPO for RLVR, requiring fewer models (no critic model) compared to the full PPO setup.
-
Training Focus
1:17:10
The model is trained only on the final response's correctness (accuracy reward). Intermediate reasoning steps are typically ignored by the reward mechanism to prevent the model from learning to obfuscate its reasoning.
-
Computational Complexity
The training process is computationally intensive, requiring significant memory (e.g., 15-20 GB RAM for small settings) and time, necessitating careful checkpointing and resource management.
Technical details
-
RLVR Workflow
3607s
The core workflow involves: 1) Generating multiple model rollouts (responses) using techniques like temperature scaling and top P filtering. 2) Computing verifiable rewards (e.g., 1 for correct answer, 0 otherwise) and format rewards (e.g., for using boxed answers). 3) Calculating relative advantages (rewards minus average/std dev). 4) Computing sequence log probabilities. 5) Calculating the final policy gradient loss (average of log probabilities weighted by advantages).
-
GRPO Loss Function
The total loss is derived from the policy gradient loss, which is calculated by averaging the sequence log probabilities weighted by the advantages. The loss function is designed to be minimized (using gradient descent) while maximizing the model's score.
-
Training Loop Implementation
The training loop follows a standard PyTorch pattern: zeroing gradients, computing the GRPO loss, performing the backward pass, and updating model weights using an optimizer (e.g., AdamW). Checkpointing is highly recommended due to the long training duration.
-
Log Probability Calculation
The speaker discusses the choice between token-level and sequence-level log probabilities for the loss calculation, noting that this choice can be treated as a hyperparameter, depending on the specific research paper or algorithm (e.g., GSPO uses sequence level, Dapo uses token level).
Mentioned resources
Channel & topics
Watch on YouTube · Back to latest
This independent, AI-assisted summary is provided for commentary and informational purposes. It may contain errors or omit important context. Please watch the original video for the creator's complete presentation. Video, thumbnail, and related copyrights belong to their respective owners.