TECHNICAL WRITING

Counting FLOPs, Params, and Memory in a Transformer

A tap-to-explain walkthrough of how FLOPs, parameters, and memory are counted in a Transformer, annotated for reading on the go. Terms underlined in sepia open a plain explanation; grids next to each formula let you tap a cell and count the multiplies and adds yourself.
expand all notes collapse all

01 Counting dots: where FLOPs come from

This started as a set of annotated notes on Google DeepMind's "How To Scale Your Model", and grew into an evergreen reference of my own, now with interactive visualizations built in.

A dot product of two length-P vectors costs 2P floating point operations: P multiplies and P adds. Below, x and y are the inputs and r is the result: tap it to see exactly which named cells of x and y it came from.

Stretch that into a matrix-vector product and you're doing one dot product per row, so an N×P matrix A against a length-P vector x costs 2NP. Tap any cell of the result y and you'll see it lights up one row of A and all of x, because that row is the only thing that fed it.

Stretch again into a full matrix-matrix product, A[N,P]·B[P,M], and now the result C is a whole grid: N rows × M cols, count them yourself: that's 2NPM FLOPs total (N×M cells, each one a 2P-FLOP dot product). Tap any cell of C and watch one row of A and one column of B light up; those are the only two things that determined it.

The general version of this rule is what actually matters for reading the accounting tables later, and it's worth being precise about what contracting and batching axes actually change, because the natural assumption is that they must affect the FLOP count differently. They don't.

Start from the matmul above: A[N,P]·B[P,M]→C[N,M], costing 2NPM. N, P, and M each appear exactly once in that formula: none of them is "extra expensive" or "extra cheap." P happens to get summed away and disappear from the output C; N and M happen to survive into it. That's the only difference between them, not how much they cost.

Now run G independent copies of that exact matmul side by side, G attention heads, say, each with its own A and B:

G independent [N,P]·[P,M] matmuls, run side by side FLOPs = G · 2NPM = 2·G·N·P·M

G shows up exactly once too, multiplying the total exactly the way P did: batching and contracting axes cost identically, per axis. The one real difference: G survives into the output (C is now shaped [G,N,M] instead of [N,M]), while P still doesn't. Whether an axis disappears from the output or not is the entire distinction between them; it has no bearing at all on how the FLOP count scales.

So the fully general rule really is as simple as it sounds: FLOPs = 2 × every axis that appears in either input, each counted once, full stop, no exception for which kind of axis it is. The categories only tell you which axes survive into the output shape:

cost of contracting C[G,H,I,J,K,L] with D[G,H,M,N,K,L] FLOPs = 2 · G·H·I·J·M·N·K·L (all eight letters, each counted once, same as always)

One consequence is easy to miss and explains a lot about why Transformers scale so well: for a square matmul, compute grows as N³ while the data you have to move only grows as N². Bigger matmuls become relatively cheaper to feed, which is exactly why architectures built almost entirely out of matmuls are the ones that scale.

compute (N³) vs. data moved (N²) as the matmul gets bigger matmul size N → ∝N³ (compute) ∝N² (data moved)
the shaded gap is arithmetic intensity: how much compute you get per byte moved (∝N³/N² = N). It's not that data movement gets cheaper as matmuls grow; compute just grows faster, so there's more work to hide that movement behind. Curves are illustrative growth rates, not to real scale.

02 Why training costs 3× a forward pass

During training you don't just run the matmul forward; you also need its gradient, and that gradient is itself two matmuls. If C = A·B, the chain rule gives you ∂L/∂B = Aᵀ·(∂L/∂C) and ∂L/∂A = (∂L/∂C)·Bᵀ, each one another 2NPM-FLOP contraction.

forward vs. training FLOPs for one matmul, PM parameters forward: 2·N·P·M (inference) backward: 4·N·P·M (one grad per input) total: 6·N·P·M (training)

That's the origin of the shorthand you'll see everywhere in scaling-law papers: training FLOPs ≈ 6 × parameters × tokens. It's not a rule about Transformers specifically; it falls straight out of how backprop works for any stack of matmuls.

WHERE THIS IS ACTUALLY WRITTEN DOWN C ≈ 6ND appears explicitly in Kaplan et al., "Scaling Laws for Neural Language Models" (2020), where "2ND for the forward pass, 4ND for backprop" is first laid out this way for language models. Hoffmann et al., "Training Compute-Optimal Large Language Models" (2022) (the Chinchilla paper) uses the same C=6ND approximation to derive its compute-optimal N/D allocation. Both treat it as an approximation, not an identity: it ignores attention's own FLOPs and a few smaller terms, which is exactly the gap §05 below quantifies.

03 Anatomy of one layer

This section walks through each block's FLOP and parameter accounting right alongside its diagram, with batch B, sequence length T, model width D, MLP width F, N query heads, K key/value heads, and head dimension H doing the bookkeeping throughout.

The MLP block

the equation, traditional form and gating variant h=σ(x⋅Win) , y=h⋅Wout
↑ traditional MLP. Two learned weight matrices, W_in and W_out, with a nonlinearity σ in between.
g=σ(x⋅Win1)⊙(x⋅Win2) , y=g⋅Wout
↑ gating variant. A third learned matrix, W_in2, projects x a second time; its output multiplies elementwise (⊙) with the first, gated branch before the same down-projection runs.
traditional: h = σ(x·W_in), y = h·W_out gating: g = σ(x·W_in1) ⊙ (x·W_in2), y = g·W_out
WHAT ABOUT THE BIAS TERM? A classic linear layer computes x·W + b, but every equation on this page, this one included, drops the bias entirely, and that's not an oversight. LLaMA, PaLM, Gemma, and most other current models set every one of these projections to have no bias at all; Chowdhery et al., "PaLM: Scaling Language Modeling with Pathways" (2022) report it also helps training stability for PaLM. Where a model does keep the bias, its cost barely registers next to the matrix it sits beside (a D-wide vector next to a D×F matrix), so it wouldn't move any of the numbers on this page either way.

In its traditional form this block is about as simple as a neural-net layer gets: project up from D to a wider hidden size F, apply a nonlinearity, then project back down to D.

Most modern implementations (LLaMA, DeepSeek, Gemma, and most other current open models) extend this with a gating einsum: a second, parallel up-projection that gets multiplied elementwise with the first before the same down-projection runs. That's 3 weight matrices of size D·F instead of the 2 above: the extra matrix buys a learned gate at the cost of 50% more MLP parameters.

Putting numbers on the gating version above (the 3-matrix one, which is what's actually running in most current models):

MLP block
operationtrain FLOPsparams
x·W_in1 (D→F)6BTDFDF
x·W_in2 (D→F)6BTDFDF
g·W_out (F→D)6BTDFDF
total≈18BTDF3DF
matrix shapesx [B,T,D] · W_in1 [D,F] → h1 [B,T,F] (W_in2 is the same shape, run in parallel) gate = h1 ⊙ h2 [B,T,F] · W_out [F,D] → y [B,T,D]

The attention block

the full equation, split into its two kinds of matmul Q=x⋅WQ , K=x⋅WK , V=x⋅WV
↑ the projections. Every W here is a learned weight matrix, fixed after training, multiplied against the activations x.
headi = softmax ( QiKiT H + mask ) ⋅ Vi
↑ dot‑product attention. No W anywhere in this line: Qi, Ki, Vi are just the per-head slices of Q, K, V computed above; both sides of every multiply here are activations, not a learned weight.
MultiHead(x) = Concat ( head1,…,headN ) ⋅ WO
↑ the 4th projection. W_O is a learned weight matrix again, the same kind of matmul as the Q/K/V lines, just on the way back out.
Q = x·W_Q, K = x·W_K, V = x·W_V ← projections (learned weights) head_i = softmax( Q_i·K_iᵀ/√H + mask ) · V_i ← dot‑product attention (no weights) MultiHead(x) = Concat(head_1, …, head_N) · W_O ← projection (learned weight)

This block runs on the same D-wide residual stream as the MLP above. Attention's whole trick is splitting that D-wide vector into several smaller "heads" it can attend with independently, then stitching them back together. That introduces a few new letters: N is the number of query heads, K is the number of key/value heads, and H is the size of each head, so N·H = D. Two more track position counts: T is the number of query positions and S is the number of key positions (usually the same sequence, so T = S).

Four weight matrices move x in and out of that per-head shape: Q, K, and V each project the D-wide input down into several H-wide heads (N of them for Q, K of them for K and V), and the output projection O puts all the heads back together into D.

It's easy to read "projections" and "dot‑product attention" as two names for the same thing: both are matmuls, and §01 already showed every matmul is nothing but dot products underneath. The real difference is what's on each side of the multiply, and it's worth being clear about before going further, because the accounting tables below treat these two as separate rows for exactly this reason:

TWO DIFFERENT KINDS OF MATMUL IN THIS BLOCK The four projections (Q, K, V, O) multiply activations by a learned weight matrix (x·W_Q, x·W_K, and so on), exactly like the MLP's matmuls. The weight side is fixed after training; only the activation side depends on how many tokens you feed in. Cost scales linearly with tokens, once per matrix.

Dot-product attention (Q·Kᵀ, then ·V) has no weight matrix anywhere in it: both sides of every multiply are activations, Q times K, then that result times V. Both sides grow with sequence length, so the cost scales with T² instead of just T. That quadratic term is the entire reason it gets its own row in the table instead of being lumped in with the projections.

Why does having two sequence-length sides turn linear scaling into quadratic? Because, as §01 already established, a matmul's cost is just (rows × cols) output cells, each one a fixed-size dot product. In a projection, only the row count (number of tokens) depends on sequence length; the column count is a fixed hyperparameter (N·H or D). In attention, both the row count (T query positions) and the column count (S key positions) depend on sequence length, so doubling the sequence length doesn't just double one side of the output grid, it doubles both sides at once, and the grid's area goes up 4×, not 2×. Scale tokens by 10× and a projection costs 10× more; attention costs 100× more.

That per-head computation, the head_i line in the equation above, is exactly what the diagram below walks through:

what one query head actually computes: head_i = softmax(Q_i·K_iᵀ/√H)·V_i Q_h [T,H] K_h [S,H] Q_h · K_hᵀ / √H → [T,S] causal mask → softmax over S V_h [S,H] · output_h [T,H]
scores are scaled by 1/√H before the softmax for numerical stability; that scaling is O(TS), negligible next to the matmul FLOPs, which is why it never shows up in the accounting tables. The amber box is where causal masking happens: future positions get set to −∞ before the softmax runs, so they contribute zero weight.

Once every head has produced its own [T,H] output, all of them get concatenated back into [T, N·H] and pass through one more projection, W_O, back down to D: the fourth matmul in the accounting table below. The dot‑product attention step shown above is the one that scales with sequence length squared rather than with model width, which is exactly why the "when does attention actually matter?" section below treats it separately from the four QKVO projection matmuls.

Putting numbers on the four projections first:

Attention projections (Q, K, V, O)
operationtrain FLOPsparams
x·W_Q (D→N·H)6BTDNHDNH
x·W_K (D→K·H)6BTDKHDKH
x·W_V (D→K·H)6BTDKHDKH
concat·W_O (N·H→D)6BTDNHDNH
total12BTD(N+K)H2D(N+K)H
matrix shapesx [B,T,D] · W_Q [D,N·H] → Q [B,T,N·H] x [B,T,D] · W_K [D,K·H] → K [B,T,K·H] (W_V is the same shape as W_K) concat(heads) [B,T,N·H] · W_O [N·H,D] → out [B,T,D]

And now the dot‑product attention step itself: the one with no weight matrix, from the diagram above:

Dot-product attention itself (batched over B and K)
operationtrain FLOPs
Q·Kᵀ6BT²NH
softmax·V6BT²NH
total≈12BT²NH
matrix shapesQ [B,N,T,H] · Kᵀ [B,K,S,H]ᵀ → scores [B,N,T,S] (each of N query heads uses its assigned K/V head, batched not contracted) scores [B,N,T,S] · V [B,K,S,H] → out [B,N,T,H]
ALSO WORTH KNOWING Whether normalization happens before or after the residual add, pre-norm vs. post-norm, doesn't change the FLOP count, but it changes training stability enough that essentially every current model (LLaMA included) has settled on pre-norm.

Attention variants: MHA, GQA, and MQA

Everything above holds for any split between N query heads and K key/value heads: that split is exactly what separates the three attention variants named everywhere.

Nothing about the computation itself changes: same equation, same per-head softmax(Q·Kᵀ/√H)·V from above, only how many distinct K and V arrays feed it. Tap through the three modes below to see the actual head-to-head assignment, not just the head counts:

query heads → key/value heads: tap a mode
MHA · N=4, K=4 GQA · N=4, K=2 MQA · N=4, K=1

Whichever mode you're in, every query head still runs its own attention computation. GQA and MQA don't reduce how many times the QKᵀ·V sequence runs, only how many distinct K and V arrays feed it, which is exactly what the accounting above already priced in via K. The real payoff shows up later, in how much has to be cached at inference time: see §08, KV caching, below.

where these were introduced[1] Vaswani et al., "Attention Is All You Need" (2017): MHA, in the original Transformer [2] Shazeer, "Fast Transformer Decoding: One Write-Head is All You Need" (2019): MQA [3] Ainslie et al., "GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints" (2023): GQA

04 Doing the accounting

You've already seen every number in this section. The MLP total (≈18BTDF) and the attention block's two totals (12BTD(N+K)H for the projections, ≈12BT²NH for dot‑product attention itself) were worked out alongside their diagrams in §03. All that's left is adding them up:

Add the MLP and the QKVO projections together (leaving dot‑product attention out for a moment) and the per-layer, per-token FLOPs collapse to something clean:

ignoring dot‑product attention itself, reasonable at short-to-medium context total training FLOPs ≈ 6 · (num tokens) · (num parameters)
how 18BTDF + 12BTD(N+K)H turns into 6 · tokens · params1) start with both totals, FLOPs and params, from §03 MLP: FLOPs = 18BTDF params = 3DF projections: FLOPs = 12BTD(N+K)H params = 2D(N+K)H 2) notice each FLOPs total is already 6 · tokens · that block's own params 18BTDF = 6 · (BT) · (3DF) 12BTD(N+K)H = 6 · (BT) · (2D(N+K)H) 3) add the two FLOPs totals and factor out the common 6·BT total FLOPs = 6·BT·(3DF) + 6·BT·(2D(N+K)H) = 6 · BT · [ 3DF + 2D(N+K)H ] 4) the bracketed term is exactly the total params per layer total params = 3DF + 2D(N+K)H (MLP params + projection params) 5) so the totals collapse to total FLOPs = 6 · BT · (total params) = 6 · (num tokens) · (num parameters)

Notice what actually did the work in step 2: it wasn't anything specific to the MLP or to attention. It's §02's fact that every projection matmul (activations times a fixed, learned weight matrix) costs 6× its own tokens×params. Add up FLOPs from any collection of projection matmuls and you're automatically adding up 6×tokens×(their combined params), regardless of which projections they are. That's the whole reason FLOPs ≈ 6·N_params·N_tokens works as a rule of thumb across entire models, not just single layers.

This is the whole justification for the rule of thumb you'll see quoted everywhere: FLOPs ≈ 6 · N_params · N_tokens. It's an approximation that quietly assumes the T²-scaling attention term is small next to everything else, which is true until it isn't.

05 When does attention actually matter?

Putting the two FLOP terms side by side (attention itself vs. everything else) and simplifying with the typical assumptions F=4D and D=NH gives a strikingly simple crossover:

fraction of matmul FLOPs spent on attention itself attention FLOPs / matmul FLOPs ≈ T / 8D
step by step, from the totals in §03/§041) start with both totals matmul FLOPs = 18BTDF + 12BTD(N+K)H (MLP + the four projections) attn itself = 12BT²NH (Q·Kᵀ and ·V, from §03) 2) apply the two simplifying assumptions F = 4D (typical MLP expansion ratio) D = NH, and N = K (standard multi-head attention, so N+K = 2N) 3) substitute and reduce each total to a single term matmul FLOPs = 18BTD(4D) + 12BTD(2NH) = 72BTD² + 12BTD(2D) (since NH = D) = 72BTD² + 24BTD² = 96BTD² attn itself = 12BT²(NH) = 12BT²D (since NH = D) 4) divide the two totals attn / matmul = 12BT²D / 96BTD² 5) cancel what's common to both: one B, one T, one D = (12/96) · (T²D)/(TD²) = (1/8) · (T/D) = T / 8D

So dot‑product attention only starts to dominate once sequence length exceeds roughly 8× the model width. For a D≈8k model that's around 64K tokens of context; for a smaller model like a 4.6k-width Gemma variant, the crossover arrives much earlier, near 37K tokens. This is why the "quadratic attention is a scaling disaster" intuition is misleading for large models but genuinely bites for small ones at long context. It's also why techniques like Flash Attention and local/sliding-window attention exist to push that crossover further out.

06 Mixture-of-Experts, briefly

An MoE layer is, to a rough first approximation, an ordinary dense layer duplicated into E independent MLP "experts," with a router sending each token to only k of them. The ratio E/k is the sparsity, typically 8–64× (DeepSeek-V3, for instance, routes each token to 8 of 256 experts).

the exact MLP numbers from §03, now with E experts and k active per token dense MLP (§03): FLOPs = 18BTDF params = 3DF MoE MLP: FLOPs = k · 18BTDF params = E · 3DF

Nothing about the block itself is new: an MoE layer really is just E copies of the exact 18BTDF / 3DF block from §03, sitting side by side. Compute scales with k because that's how many of those copies any single token actually runs through; parameters scale with E because all of them have to be stored whether or not a given token uses them. Divide one scaling by the other and you get E/k back: DeepSeek-V3's E=256, k=8 means 256× the parameter capacity of a single dense MLP for only 8× its compute.

That buys you a much larger parameter count for a nearly fixed compute budget per token. The bill comes due in communication instead of FLOPs: routing tokens to their assigned experts and back requires an extra pair of all-to-all exchanges per MoE layer that a dense model never needs.

07 Gradient checkpointing (rematerialization)

Backprop needs the intermediate activations from the forward pass to compute gradients: every input to every nonlinearity, in principle. Saving all of them is what turns backprop's compute cost from quadratic-in-layers into linear-in-layers, but the memory bill is brutal: a model with 4M tokens per batch, 64 layers, and D=8192 would need on the order of ~84TB of activations in bf16 if it saved everything.

WHERE THE "~20 TENSORS PER LAYER" NUMBER COMES FROM It's a rough tally of every intermediate node in the layer diagram that a naive autodiff graph would otherwise keep alive: both projections' inputs, both attention matmuls' operands, every nonlinearity's pre-activation and post-activation value (you need g(x) and exp(g(x)) to differentiate exp(g(x)), for example). None of it is optional under naive autodiff; rematerialization is entirely about choosing which of those ~20 you're willing to recompute instead of store.

Two policies bracket the tradeoff space:

Block remat: save only each layer's input, nothing else. Most aggressive, cuts memory to roughly 1/20th, but the backward pass has to redo essentially the entire forward computation to reconstruct everything else, pushing training FLOPs from the usual 6·params·tokens up to roughly 8·params·tokens. That's a real compute tax paid specifically for the memory you saved.

Big-matmuls-only: save just the outputs of the large matmuls (the ones that are actually expensive to recompute) and let cheaper activation functions and parts of attention be recomputed. This lands around 7 saved tensors per layer instead of 20, at a much smaller compute penalty than block remat, because what's cheap to recompute (elementwise ops) gets recomputed and what's expensive (matmuls) gets kept.

WHY THIS MATTERS BEYOND THE FLOP COUNT The real design problem is rarely "minimize FLOPs" in isolation; it's picking a remat policy against a fixed HBM budget, batch size, and pipeline/sharding scheme simultaneously, since the "extra" recompute FLOPs and the activation memory footprint trade off directly against each other on the same roofline. In JAX this is controlled per-operation with jax.remat / jax.checkpoint, which lets you set policy at the granularity of individual ops rather than only per-layer.

08 KV caching

Inference splits into two phases with very different cost profiles: prefill, which processes the prompt in one shot and stores the resulting key/value projections, and generation, which reuses that cache to sample one token at a time without recomputing attention over the whole prefix.

size of one KV cache, int8, "2" for keys+values bytes = 2 · S · L · K · H

Plug in realistic numbers (8K context, 64 layers, K·H=D=8192) and you land around 8 GiB per sequence. That single formula is the entire reason grouped-query attention exists: K (the number of KV heads) is the one lever in this equation that shrinks the cache without touching model width D, which is why production inference systems push K far below N.

09 Takeaways

MLP dominates params and FLOPs: 3DF parameters vs. attention's 2D(N+K)H, and the same holds for FLOPs, as long as T stays under ~8D.
6·params·tokens is a good training-FLOPs estimate for dense models at moderate context length; it's what falls out of forward (2×) + backward (4×) for every matmul in the network.
KV cache size is roughly 2·S·L·K·H per sequence; the single biggest lever on it is shrinking K via multi-query or grouped-query attention.
Rematerialization is a knob, not a fixed cost: anywhere from ~6× to ~8× params·tokens in training FLOPs depending on how aggressively you trade recompute for activation memory.