ReLU² Activation Function in NanoChat: Implementation and Usage Guide
The ReLU² activation function in NanoChat is mathematically defined as (max(0, x))² and implemented in the MLP class as F.relu(x).square(), providing a computationally efficient non-linearity that amplifies positive signals more aggressively than standard ReLU.
The NanoChat repository by Andrej Karpathy demonstrates a minimal GPT implementation with unique architectural choices. Inside nanochat/gpt.py, the feed-forward network utilizes ReLU², which applies element-wise squaring after the rectified linear unit to create a distinct activation curve optimized for transformer training dynamics.
What Is the ReLU² Activation Function?
Unlike standard ReLU (Rectified Linear Unit) which outputs max(0, x), or GELU used in many modern transformers, ReLU² first clips negative values to zero then squares the positive remaining values. This produces the mathematical operation:
output = torch.where(x > 0, x ** 2, 0)
In PyTorch functional notation, this becomes F.relu(x).square(). The squaring operation increases the magnitude of larger activations exponentially while maintaining the sparsity benefits of ReLU, as negative inputs remain exactly zero.
Source Code Implementation in NanoChat
The specific implementation resides in the MLP class within the repository's core model file. A comment on lines 35-38 explicitly notes the design choice as "relu^2 activation in MLP", guiding developers to the relevant architectural section.
The actual computation occurs in the forward pass, where the code on lines 136-138 chains the operations:
def forward(self, x):
x = self.c_fc(x)
x = F.relu(x).square() # ReLU squared activation
x = self.c_proj(x)
return x
Here, c_fc represents the first linear projection (typically expanding dimensionality by 4×), followed by the ReLU² activation, then c_proj projects back to the model dimension. This placement follows the standard transformer MLP pattern but substitutes GELU or SiLU with the squared ReLU variant.
Advantages of ReLU² in Transformer Architectures
Several characteristics make ReLU² suitable for this minimal GPT implementation:
- Stronger non-linearity: Squaring positive values creates a quadratic response curve, allowing the network to learn higher-order feature interactions without additional parameters.
- Computational efficiency: Compared to GELU or Swish activations requiring exponential or sigmoid calculations,
F.relu(x).square()executes as two simple tensor operations optimized by PyTorch's backend. - Gradient flow: The derivative
2*xfor positive inputs provides linear gradient scaling, avoiding the vanishing gradient problems found in saturating activations while maintaining sparsity for negative values. - Simplicity: The implementation requires no custom autograd functions or complex branching, aligning with NanoChat's philosophy of minimal, readable code.
Practical Implementation Examples
To implement ReLU² independently or verify its behavior:
import torch
import torch.nn.functional as F
def relu_squared(x: torch.Tensor) -> torch.Tensor:
"""
Implements the ReLU² activation function as used in NanoChat.
Equivalent to (max(0, x)) ** 2
"""
return F.relu(x).square()
# Example tensor with positive and negative values
test_input = torch.tensor([-2.0, -1.0, 0.0, 1.0, 2.0])
activated = relu_squared(test_input)
print(activated) # tensor([0., 0., 0., 1., 4.])
When integrating into a custom MLP module mimicking NanoChat's architecture:
import torch.nn as nn
import torch.nn.functional as F
class NanoChatMLP(nn.Module):
def __init__(self, n_embd: int):
super().__init__()
self.c_fc = nn.Linear(n_embd, 4 * n_embd)
self.c_proj = nn.Linear(4 * n_embd, n_embd)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.c_fc(x)
x = F.relu(x).square() # ReLU² activation
x = self.c_proj(x)
return x
This matches the architectural pattern found in nanochat/gpt.py, where the expansion factor of 4 is standard for transformer feed-forward networks.
Summary
The ReLU² activation in NanoChat provides a distinctive approach to non-linearity in transformer MLP blocks:
- Defined mathematically as (max(0, x))² and implemented via
F.relu(x).square() - Located in
nanochat/gpt.pywithin the MLP class (lines 136-138) with explicit documentation (lines 35-38) - Offers computational efficiency over GELU while maintaining strong gradient flow
- Serves as an example of how minimal architectural variations can achieve competitive performance in small-scale language models
Frequently Asked Questions
How does ReLU² differ from GELU activation?
GELU (Gaussian Error Linear Unit) applies a probabilistic gate using the cumulative distribution function of the standard normal distribution, creating smooth transitions around zero. In contrast, ReLU² maintains a hard zero threshold like standard ReLU but squares positive values, creating a sharp quadratic edge at zero. GELU requires more computation (error function calculations), whereas F.relu(x).square() executes faster on modern hardware while still providing non-linear capacity.
Can I replace ReLU² with standard ReLU in NanoChat?
Yes, you can substitute F.relu(x).square() with F.relu(x) or F.gelu(x) in the forward method of the MLP class, but this changes the model's learning dynamics. Standard ReLU provides linear scaling for positive values, which may reduce the network's ability to model complex interactions in the expanded MLP dimension, potentially requiring training adjustments or architecture modifications to maintain convergence.
Why use square() instead of pow(2) in the implementation?
The .square() method in PyTorch is a specialized operation optimized specifically for squaring tensors, often implemented with faster kernels than the general-purpose .pow(2) function. While mathematically equivalent, F.relu(x).square() follows PyTorch best practices for numerical stability and performance, aligning with NanoChat's goal of efficient, clean implementation as seen in the source code.
Is ReLU² used in both MLP layers and attention mechanisms?
In the NanoChat implementation, ReLU² appears exclusively in the MLP (feed-forward) layers of the transformer blocks. The attention mechanism typically uses softmax-based normalization and does not apply activation functions to query/key/value projections. The squared activation specifically processes the intermediate representations between the c_fc and c_proj linear transformations within the MLP submodule.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →