← All Posts
From Scratch · C++ · Transformers · Attention Variants · Part 2 of 13

Softmax Attention: Numerically Stable Softmax

Introduction

Softmax is the function that turns attention scores into weights. At first glance, it seems simple: $\text{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}$. But computing this naively leads to overflow/underflow: if any input is large, $e^{x}$ explodes; if all inputs are small, $e^{x}$ underflows to zero.

This post breaks down the numerically stable version: subtract the maximum before exponentiating, sum the exponents, and divide. We'll implement it in C++, compare it to torch::softmax, and trace the three steps with an interactive animation.

Why Naive Softmax Fails

The textbook definition:

softmax(x_i) = exp(x_i) / sum_j(exp(x_j))

The problem: if $x = [1000, 1001, 999]$, then $e^{1000}$ is beyond float64 range. The function overflows and returns NaN. Even if the inputs fit, the ratio can be unstable due to floating-point rounding errors.

The Numerically Stable Trick

Use the log-sum-exp trick: subtract the maximum before exponentiating.

Let c = max(x) softmax(x_i) = exp(x_i - c) / sum_j(exp(x_j - c)) Numerically equivalent, but safe: - Largest term becomes exp(0) = 1 - All other terms are exp(negative), avoiding overflow - Floating-point precision is preserved

C++ LibTorch Implementation

struct SoftmaxAttentionImpl : public torch::nn::Module { torch::nn::Linear wq{nullptr}, wk{nullptr}, wv{nullptr}; int64_t d_model, d_k; SoftmaxAttentionImpl(int64_t d_model_) : d_model(d_model_), d_k(d_model_) { wq = register_module("wq", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wk = register_module("wk", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wv = register_module("wv", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); } torch::Tensor stable_softmax(torch::Tensor x) { // Step 1: Subtract max for numerical stability auto x_max = torch::max(x, -1, true).values; auto x_shifted = x - x_max; // Step 2: Compute exp and sum auto exp_x = torch::exp(x_shifted); auto sum_exp = torch::sum(exp_x, -1, true); // Step 3: Divide by sum return exp_x / sum_exp; } torch::Tensor forward(torch::Tensor x) { auto Q = wq(x); auto K = wk(x); auto V = wv(x); auto scores = torch::matmul(Q, K.transpose(-2, -1)); scores = scores / sqrt((float)d_k); // Use our stable softmax instead of torch::softmax auto attn_weights = stable_softmax(scores); auto output = torch::matmul(attn_weights, V); return output; } }; TORCH_MODULE(SoftmaxAttention);

The Three Steps Explained

StepOperationLibTorch CodePurpose
1. Shift$x_i - \max(x)$torch::max(x, -1, true).valuesSubtract max to avoid overflow
2. Exp & Sum$e^{x_i - c}$ and $\sum e^{x_j - c}$torch::exp(x_shifted) + torch::sum(..., -1, true)Compute exponents and their sum
3. Divide$\frac{e^{x_i - c}}{\sum e^{x_j - c}}$exp_x / sum_expNormalize to get probabilities

In the code, the second argument to torch::max and torch::sum is the dimension (-1 for last), and the third argument is true for keepdim. This preserves the shape: if input is (batch, T, T), max returns (batch, T, 1), which broadcasts correctly when we subtract.

Stable vs. torch::softmax

LibTorch's built-in torch::softmax already uses numerical stabilization internally. For production code, prefer it. But understanding the manual implementation is valuable:

// Both are equivalent and numerically stable: // Option 1: Use torch::softmax (recommended) auto weights = torch::softmax(scores, -1); // Option 2: Manual stable softmax (educational) auto x_max = torch::max(scores, -1, true).values; auto exp_x = torch::exp(scores - x_max); auto sum_exp = torch::sum(exp_x, -1, true); auto weights = exp_x / sum_exp;

The manual version is slower (four separate operations instead of one fused kernel), but it shows the underlying mathematics. Most modern frameworks fuse these operations for performance.

Interactive Softmax Steps

▶ Three Steps of Numerically Stable Softmax

Watch how subtraction of max, exponentiation, and normalization work on a small logit vector.

Complete Shape Trace

OperationShape
scores input(batch, T, T)
max(scores, dim=-1, keepdim=true)(batch, T, 1)
scores - max(batch, T, T) [broadcasted]
exp(shifted)(batch, T, T)
sum(exp, dim=-1, keepdim=true)(batch, T, 1)
exp / sum [broadcasted](batch, T, T)

Key Insights

Numerical stability is not optional. Even though the math is simple, naive implementations fail on real data. The max-subtract trick is one of the most important pieces of lore in deep learning.
Keepdim=true is essential. When we compute max or sum along a dimension, we need to keep that dimension (even if size 1) so that broadcasting works when we subtract or divide.
Softmax converts scores to probabilities. After softmax, each row of the attention matrix is a probability distribution: non-negative, sums to 1, interpretable as "how much weight each token gets."

What Comes Next

Self-attention computes attention across all tokens. But in autoregressive models (language models, text generation), a token should never see future tokens. This is enforced by a causal mask that sets future scores to $-\infty$ before softmax. Next: Causal Attention.

Previous Self-Attention