How KL Divergence Prevents Mode Collapse in RLHF Training

KL divergence prevents mode collapse by acting as a soft regularizer that penalizes large deviations from a frozen reference model's distribution, ensuring the policy maintains diversity while optimizing for human feedback rewards.

Reinforcement Learning from Human Feedback (RLHF) fine-tunes language models using Proximal Policy Optimization (PPO), but unconstrained reward optimization often causes the policy to collapse into a narrow set of high-reward outputs. The Lordog/dive-into-llms repository demonstrates how integrating a KL-divergence penalty between the active policy and a reference model maintains distributional diversity throughout training.

Understanding Mode Collapse in RLHF

Mode collapse occurs when a policy optimized purely on external rewards converges to repetitive, high-scoring outputs while ignoring the broader distribution of plausible responses. In the RLHF training loop implemented in documents/chapter11/RLHF.ipynb, the PPO algorithm drives updates based on feedback from a reward model (such as a sentiment classifier). Without constraints, the policy quickly discovers and over-optimizes for specific patterns that maximize the reward signal, resulting in degenerate output distributions.

The repository addresses this by treating the original pretrained model as a reference distribution that the fine-tuned policy must not deviate from too substantially. This is achieved by loading two copies of the base model: one that gets updated (model) and one that remains frozen (ref_model), as instantiated on line 222 of the notebook.

How KL Divergence Acts as a Regularizer

KL divergence serves as an additional reward term that penalizes the policy for moving too far from the reference model's behavior. According to the repository's documentation in documents/chapter11/README.md (line 14), this regularization technique is explicitly designed to prevent the optimization process from collapsing to a single mode.

Mathematical Formulation

The implementation computes the KL divergence using the standard formula:

$$ \text{KL} = \sum p_{\text{model}} \cdot \log\left(\frac{p_{\text{model}}}{p_{\text{ref}}}\right) $$

This value is added to the scalar reward returned by the human-feedback model. The PPO optimizer then maximizes the combined objective: reward + λ·(−KL), where the coefficient λ controls the regularization strength. Line 15 of documents/chapter11/RLHF.ipynb explicitly references this KL-divergence penalty as an additional reward component.

Architectural Flow

The training loop follows this specific sequence to implement the regularization:

  1. Load dual models: Initialize both the trainable policy and the frozen reference model from the same base weights.
  2. Generate responses: Sample outputs using the current policy.
  3. Compute log-probabilities: Calculate the likelihood of generated tokens under both the active policy and the reference model.
  4. Calculate KL penalty: Compute the divergence between the two distributions and subtract it from the external reward.
  5. PPO update: Apply gradients to maximize the penalized reward, keeping the policy close to the reference distribution.

Implementation in the dive-into-llms Repository

The following excerpt from documents/chapter11/RLHF.ipynb demonstrates the practical implementation using the trl library:

from transformers import AutoModelForCausalLMWithValueHead, AutoTokenizer
from trl import PPOTrainer, PPOConfig

# 1️⃣ Load policy and frozen reference model

config = PPOConfig(model_name="model/gpt2-imdb", learning_rate=1.41e-5, log_with="wandb")
policy = AutoModelForCausalLMWithValueHead.from_pretrained(config.model_name)
ref_policy = AutoModelForCausalLMWithValueHead.from_pretrained(config.model_name)
tokenizer = AutoTokenizer.from_pretrained(config.model_name)
tokenizer.pad_token = tokenizer.eos_token

# 2️⃣ Initialise PPO trainer (handles KL automatically)

ppo_trainer = PPOTrainer(config, policy, ref_policy, tokenizer, dataset=your_dataset)

# 3️⃣ Inside the training loop

for batch in ppo_trainer.dataloader:
    # generate a response with the current policy

    query_tensors = batch["input_ids"]
    response_tensors = ppo_trainer.generate(query_tensors, **gen_kwargs)

    # compute external reward (e.g., sentiment score)

    texts = [q + r for q, r in zip(batch["query"], response_tensors)]
    reward = sentiment_pipe(texts)[...]   # positive‑score tensor list

    # 4️⃣ PPO step automatically adds KL penalty

    stats = ppo_trainer.step(query_tensors, response_tensors, reward)
    ppo_trainer.log_stats(stats, batch, reward)

Key implementation details from the source code:

  • ref_policy remains frozen throughout training, providing the baseline distribution for divergence calculations.
  • PPOTrainer.step internally computes kl = kl_divergence(policy, ref_policy) and applies the penalty automatically.
  • The divergence is calculated per-token and aggregated across sequences before being subtracted from the batch rewards.

Tuning the KL Coefficient for Optimal Results

The strength of the regularization is controlled by config.kl_coef, which defaults to 0.2 in the PPOConfig. Adjusting this coefficient allows you to balance between two competing objectives:

  • Higher values (e.g., 0.5-1.0): Force the policy to stay closer to the original model, preserving diversity but potentially limiting reward optimization.
  • Lower values (e.g., 0.1-0.2): Allow more aggressive optimization for human feedback, increasing the risk of mode collapse but achieving higher external rewards.

As implemented in the dive-into-llms codebase, maintaining a modest coefficient (around 0.2) provides sufficient regularization to prevent collapse while still allowing meaningful adaptation to the reward model.

Summary

  • KL divergence acts as a soft regularizer by penalizing deviations from the reference model's token distribution, preventing the policy from collapsing into narrow, high-reward modes.
  • Dual model architecture in documents/chapter11/RLHF.ipynb maintains a frozen reference alongside the trainable policy to compute the divergence term.
  • Automatic integration via trl.PPOTrainer handles the KL calculation within the training step, applying the penalty as described on line 15 of the notebook.
  • Coefficient tuning via config.kl_coef (default 0.2) provides control over the exploration-exploitation trade-off during RLHF optimization.

Frequently Asked Questions

What exactly is mode collapse in RLHF training?

Mode collapse occurs when a language model policy optimized via PPO discovers and exclusively produces a narrow subset of outputs that maximize the reward model's score, ignoring the broader distribution of valid responses. This results in repetitive, degenerate text generation that fails to capture the diversity of human language, effectively "collapsing" the model's distribution to a few high-reward modes.

Where is the KL penalty calculated in the dive-into-llms codebase?

The KL-divergence penalty is described conceptually on line 14 of documents/chapter11/README.md and implemented practically on line 15 of documents/chapter11/RLHF.ipynb. In the code, trl.PPOTrainer.step() handles the internal calculation by comparing log-probabilities between the active policy and the frozen ref_policy (instantiated on line 222), then subtracting the scaled divergence from the external reward signal.

How does the KL coefficient affect model performance?

The kl_coef parameter (default 0.2 in PPOConfig) determines the weight of the divergence penalty in the optimization objective. Higher values force the policy to remain closer to the reference model, preserving linguistic diversity but potentially limiting how much the model can adapt to human feedback. Lower values permit greater optimization for external rewards but increase the risk of mode collapse and unstable training dynamics.

Can KL divergence completely eliminate mode collapse?

While KL divergence significantly mitigates mode collapse by maintaining distributional alignment with the reference model, it does not guarantee complete elimination. If the reward model itself has blind spots or if the KL coefficient is set too low, the policy may still exploit loopholes in the reward function. The technique is most effective when combined with careful reward modeling and appropriate coefficient tuning.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →