Value Head Architecture in RLHF Training: A Deep Dive into the Critic Network

The RLHF value head is a two-layer MLP with ReLU activation that projects transformer hidden states to scalar value estimates, zero-initialized to prevent early training instability.

In Reinforcement Learning from Human Feedback (RLHF), the critic network estimates state values to guide policy optimization. The train-llm-from-scratch repository implements this critical component as a lightweight scalar head attached to the transformer backbone, defined in src/post_training/value_head.py.

Value Head Architecture Overview

The value head sits atop the transformer backbone, consuming token-level hidden representations and producing a single scalar per token. This scalar serves as the critic's estimate $V(s_t)$ for Proximal Policy Optimization (PPO).

Two-Layer MLP Structure

According to the source code in src/post_training/value_head.py (lines 30-34), the architecture consists of:

  1. Input projection: nn.Linear mapping from $n_{\text{embed}}$ to $n_{\text{embed}}$
  2. Non-linearity: nn.ReLU activation
  3. Output projection: nn.Linear mapping from $n_{\text{embed}}$ to $1$

The hidden dimension $n_{\text{embed}}$ is sourced from the transformer's language model head (lm_head.in_features). After processing, the output is squeezed to shape (B, T), representing batch size and sequence length.

Zero Initialization for Training Stability

Lines 36-38 of src/post_training/value_head.py implement zero initialization for the final linear layer's weights and bias. This ensures the critic starts with near-zero value predictions, preventing the value network from destabilizing the policy during early training phases before the backbone has learned useful representations.

Integration with the Transformer

The TransformerWithValueHead class wraps the base transformer to provide both policy logits and value estimates.

The TransformerWithValueHead Wrapper

This wrapper (defined in src/post_training/value_head.py) extends the base transformer with two key capabilities:

  • Dual output: Returns both policy logits (from lm_head) and value estimates (from the value head)
  • Context forwarding: Exposes the backbone's context_length property for rollout utilities

Forward Pass Methods

The implementation provides two inference modes:

forward(idx): Computes hidden states through the transformer backbone, then branches to produce both lm_head logits and value head scalars. Returns a tuple (logits, values) where logits have shape (B, T, vocab_size) and values have shape (B, T).

value_only(idx): A torch.no_grad efficient path that returns only the per-token values without computing the expensive vocabulary-sized logits. This method is essential when scoring rollouts during PPO training.

Code Examples

Building a PPO Actor-Critic


# Example: building a PPO actor‑critic from a pretrained transformer

from src.models.transformer import Transformer          # the language model backbone

from src.post_training.value_head import TransformerWithValueHead

# 1️⃣ Load or instantiate the transformer backbone

backbone = Transformer.from_pretrained("gpt2")   # any compatible model

# 2️⃣ Wrap it with the value head

actor = TransformerWithValueHead(backbone).to("cuda")

# 3️⃣ Forward pass – get policy logits and value estimates

tokens = torch.arange(0, 10).unsqueeze(0).to("cuda")   # shape (1, seq_len)

logits, values = actor(tokens)                        # logits: (1, seq_len, vocab), values: (1, seq_len)

# 4️⃣ Scoring rollouts (no logits needed)

rollout_values = actor.value_only(tokens)             # shape (1, seq_len)

Using the Value Head in PPO Training


# Example: using the value head in a PPO training loop (simplified)

for batch in dataloader:
    idx = batch["input_ids"].to(device)

    # Get policy and value predictions

    logits, values = actor(idx)

    # Compute log‑probs, advantages, PPO loss, etc.

    # (the actual PPO implementation lives in src/post_training/ppo.py)

    loss = ppo_loss_fn(logits, values, batch["rewards"], ...)
    loss.backward()
    optimizer.step()

Summary

  • The value head is a two-layer MLP ($n_{\text{embed}} \to n_{\text{embed}} \to 1$) with ReLU activation
  • Zero initialization of the final layer (lines 36-38) ensures stable early training
  • TransformerWithValueHead wraps the base model to provide both policy and value outputs
  • The value_only() method optimizes inference by skipping logits computation during rollout scoring
  • Located in src/post_training/value_head.py, integrated with PPO training in src/post_training/ppo.py

Frequently Asked Questions

Why is the value head zero-initialized in RLHF training?

Zero initialization prevents the critic from outputting extreme value estimates at the start of training. According to the implementation in lines 36-38 of src/post_training/value_head.py, setting both weights and bias to zero ensures the critic starts with neutral predictions, allowing the policy to explore without interference from an over-confident value network before the transformer backbone has learned meaningful representations.

What is the difference between forward() and value_only() in the value head implementation?

The forward() method computes both policy logits and value estimates through the full transformer and both heads, returning a tuple (logits, values). The value_only() method uses torch.no_grad context and bypasses the lm_head computation entirely, returning only the scalar value estimates. This distinction is critical for PPO efficiency, as value_only() avoids the expensive vocabulary projection when you only need state values for advantage calculation during rollout generation.

How does the value head architecture handle different model sizes?

The architecture automatically adapts to the backbone's embedding dimension by reading lm_head.in_features during initialization. Whether wrapping a small GPT-2 or a larger transformer variant, the value head scales its hidden layers to match the model's $n_{\text{embed}}$ dimension, ensuring consistent interface compatibility across different model configurations in the train-llm-from-scratch framework.

What shape does the value head output and how is it used in PPO?

The value head outputs a tensor of shape (batch_size, sequence_length) or (B, T), representing a scalar value estimate $V(s_t)$ for each token position. In PPO training (implemented in src/post_training/ppo.py), these per-token values are used to compute advantages and returns, comparing predicted values against actual discounted rewards to update the critic network.

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 →