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

> Explore the RLHF value head architecture: a two-layer MLP projecting transformer states to scalar values. Learn how zero-initialization ensures stable training. Understand the critic network in LLM training from scratch.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: deep-dive
- Published: 2026-06-11

---

**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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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

```python

# 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

```python

# 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/value_head.py), integrated with PPO training in [`src/post_training/ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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.