Key Concepts
- Reinforcement Learning for Language Models (RL for LMs): Using RL to fine-tune LMs.
- States, Actions, Rewards: Defining these elements for LMs. State = prompt + response. Action = next token. Reward = how good the response is.
- Outcome Rewards vs. Process Rewards: Outcome rewards are based on the entire response; process rewards are given during generation.
- Verifiable Rewards: Deterministic reward functions.
- Policy Gradient: A class of methods for learning policies using gradient methods.
- Baseline: A function B(S) that depends only on the state (S) and is subtracted from the reward to reduce variance in policy gradient updates.
- Advantage Function: Measures how much better an action is compared to the expected value from a state following the current policy.
- GRPO (Generalized Policy Region Optimization): A policy optimization algorithm that simplifies PPO (Proximal Policy Optimization) in the context of language models by leveraging the natural group structure of responses from a single prompt.
- KL Penalty: An additional regularization term that penalizes deviations from a reference policy.
- Freezing Parameters: Treating certain parameters or models as constant during training to prevent unintended gradient flow.
Reinforcement Learning Setup for Language Models
The lecture begins by clarifying the reinforcement learning setup when applied to language models:
- State: Defined as the concatenation of the prompt and the response generated so far. Each state represents a sequence of tokens.
- Action: The generation of a single token, which is appended to the state.
- Reward: A function of the entire generated response, focusing on outcome rewards that are verifiable (deterministic computation). This avoids the need for human evaluation.
- Transition Dynamics: Simple concatenation of the action (token) to the state. This determinism is a significant advantage compared to general RL, allowing for planning.
The lecture emphasizes that in language models, the notion of "state" is flexible because it's based on tokens and the model can create its own "scratch pad." The challenge is ensuring that these tokens lead to a correct ground truth answer.
Policy Gradient Methods
The core of the lecture dives into policy gradient methods for optimizing the language model policy.
-
Objective: Maximize the expected reward with respect to the policy π.
-
Policy Gradient Theorem: The gradient of the expected reward can be expressed as:
∇ J(π) = E[∇log π(a|s) * R(s, a)]
Where: * J(π) is the expected reward * π(a|s) is the policy (probability of action a given state s) * R(s, a) is the reward for taking action a in state s.
-
Naive Policy Gradient: Sample prompts and responses, then update parameters based on the sampled gradient. It is analogous to SFT, but weighted by the reward.
- If the reward is binary (0 or 1), updates only occur for correct responses. This can lead to issues with sparse rewards where a bad policy gets stuck because it rarely receives positive feedback.
Baselines for Variance Reduction
The lecture addresses the high variance inherent in policy gradient methods.
-
Baseline Concept: Subtracting a baseline function B(S) from the reward to reduce variance without changing the expected reward:
E[R(s, a) - B(s)]
- B(S) must depend only on the state S and not on the action A.
-
Toy Example: A two-state example illustrates how a poorly chosen action in an "easier" state (higher reward) can be incorrectly reinforced without a baseline.
-
Optimal Baseline: The optimal baseline minimizes variance and has a closed-form solution, but it's generally difficult to compute.
B*(s) = E[(∇log π(a|s))^2 * R(s, a)] / E[(∇log π(a|s))^2]
-
Heuristic Baseline: A practical approximation sets the baseline to the expected reward given the state:
B(s) = E[R(s, a) | s]
-
This heuristic connects to the concept of the advantage function:
A(s, a) = Q(s, a) - V(s)
Where: * Q(s, a) is the Q-function (expected reward given state s and action a). * V(s) is the value function (expected reward given state s).
-
Using the heuristic baseline is equivalent to optimizing the advantage function.
-
-
General Form of Policy Gradient Updates:
Update ∝ ∇log π(a|s) * δ
Where δ is some estimate based on the reward. It could be: * δ = R (Naive Policy Gradient) * δ = R - B(s) (Baseline Subtraction) * δ = (R - B(s)) / σ (GRPO - Normalizing by Standard Deviation).
GRPO: Deep Dive
The lecture then focuses on GRPO and provides a detailed explanation with code examples.
- Motivation for GRPO: Language models allow generating multiple responses from a single prompt, creating a natural "group structure" for comparison, which motivates GRPO.
GRPO Pseudo Code (simplified)
- Generate Responses: Given a prompt, generate N responses using the language model policy.
- Compute Rewards: Calculate the reward for each of the N responses.
- Compute Deltas: Calculate the delta values using rewards. Common choices are:
- Rewards (Naive policy gradient).
- Centered Rewards (Subtract the mean reward of the group).
- Normalized Rewards (Center and divide by the standard deviation of the group).
- Compute Log Probabilities: Calculate the log probabilities of the actions (tokens) in each response under the current policy.
- Compute GRPO Loss:
- Calculate the ratio of the log probability of each action under the current policy to the log probability of the same action under a frozen, old policy (importance weighting).
- Clip these ratios to a predefined range (e.g., 1-ε to 1+ε) to prevent excessively large policy updates.
- Multiply the clipped ratios by the delta values and take the mean.
- KL Penalty (Optional): Add a penalty to the loss to regularize the policy and prevent it from deviating too far from a reference policy.
- Update Policy: Update the policy parameters using the calculated loss.
Example: Sorting Task
A simple sorting task is used to illustrate the algorithm and code.
- Task: Sort a sequence of N numbers.
- Prompt: A list of N unsorted numbers.
- Response: A list of N sorted numbers.
- Reward Function (Examples):
- Binary (1 if sorted, 0 otherwise) - leads to sparse reward.
- Number of positions matching the ground truth (gives partial credit).
- Combination of token inclusion and adjacent pair sorting (gives more partial credit, but can have loopholes).
Model
A simple, non-autoregressive model is defined.
- Fixed Input/Output Length: The prompt and response lengths are fixed to simplify the code.
- Positional Information: Per-position parameters are used to capture positional information.
- Independent Decoding: Each position in the response is decoded independently (non-autoregressive).
Code Walkthrough
The lecture then walks through the code, covering:
- Generating Responses: Sampling responses from the model's logits.
- Computing Rewards: Calculating the reward for each generated response using a defined reward function.
- Computing Deltas: Transforming rewards into delta values for policy updates (centering, normalization).
- Computing Log Probabilities: Calculating the log probabilities of the generated tokens under the current policy.
- Computing the Loss: Implementing the GRPO loss function, including clipping.
- KL Penalty: Implementing the KL divergence penalty for regularization.
- Training Loop: Combining all components into a training loop with outer epochs and inner gradient steps.
Implementation Details
Several important implementation details are highlighted:
- Freezing Parameters: Using
torch.no_grad()to freeze parameters of the old policy or reference model during loss calculation. - Multiple Models: Maintaining the current policy, the old policy (for GRPO), and potentially a reference policy (for KL regularization).
- Inference Cost: Inference can be expensive, hence the inner loop performs multiple gradient steps on the same set of generated responses.
Experiments
The lecture presents experimental results on the sorting task, comparing different reward functions and delta computation methods.
- Observations:
- Partial credit reward functions can lead to getting stuck in local optima.
- Centering rewards can improve performance by pushing the model away from incorrect responses.
- The loss function may not be a reliable indicator of performance because the distribution of responses changes over time.
Conclusion
The lecture concludes by emphasizing:
- The importance of reinforcement learning: To surpass human abilities by optimizing towards measurable rewards.
- The challenge of reward design: Creating rewards that are not easily "hacked" and generalize well.
- The complexity of building RL systems: Scaling RL for LMs requires managing multiple models, distributed inference, and parallel processing.
The lecture highlights that while policy gradient frameworks are conceptually clear, building and scaling RL systems for language models introduces significant engineering challenges beyond pre-training.
AI summaries can miss context or contain errors. Check important details against the original video.





