A speech encoder's normalization, activation and residual-add chain dominates step time — what does fusing it into one pass remove?
answer
- the arithmetic never changes
- count round trips to device memory
- six tensor transfers where two would do
- only helps left of the ridge point
- gain capped by the chain's share of the step
basics
~20 sFusion removes the trips to device memory between the stages. Unfused, the activation tensor is written and re-read after each operation; fused, it is read once and written once. The arithmetic is identical — only the memory traffic shrinks.
solid answer
~50 sEach stage in that chain is a few FLOPs per element, so all three are bandwidth-bound: their cost is the traffic, not the maths. Run separately, the tensor makes about three round trips through device memory — a read and a write per stage — plus per-launch overhead and temporary buffers. Fused into a single pass, the data is read once, normalized, passed through the activation and added to the residual while it is still on-chip, and written once, so the chain's time falls toward a third of what it was. Fusion does not reduce FLOPs, does not help a compute-bound matmul, and cannot save more than the chain's share of step time. It is also constrained: the mean and variance are a reduction across the feature dimension, so that reduction must complete before the scaling, and the backward pass has its own traffic to fuse separately.
go deeper
Be ready to say that running three small operations one after another makes the data travel to and from memory three times, and that doing them in a single pass over the data is faster even though the maths is identical.
Explain the traffic arithmetic: read plus write per stage unfused versus one read and one write fused, and why that ratio is the speedup for operations whose intensity is under one FLOP per byte.
Demonstrate scoping. Measure the chain's share of step time, state the ceiling the fusion can reach, and say explicitly which parts of the step it cannot touch — the matmuls, the backward traffic, an idle device.
Own the tradeoff between hand-shaped fused paths and keeping the model easy to change. Fused chains are faster but harder to modify and to reason about numerically, so argue when the throughput win justifies the maintenance cost and when it does not.
## Why this chain is a bandwidth problem All three stages do trivial amounts of arithmetic per element. A normalization over the feature dimension computes a mean and a variance across each token's feature vector, then applies `(x - mean)/sqrt(var + eps)` with a learned scale and shift — a small constant of FLOPs per element. A smooth activation is one value in, one value out. A residual add is one FLOP per element. Their arithmetic intensity is well under one FLOP per byte, which puts every one of them on the bandwidth-limited slope of the roofline. Their runtime is essentially `bytes moved / achievable bandwidth`. ## What running them separately actually costs Suppose the activation tensor is `B` bytes. Executed as three independent passes, each one reads its input from device memory and writes its output back: ``` normalize: read B, write B activate: read B, write B add: read B (+ read the residual), write B ``` That is roughly six tensor-sized transfers where the chain only *needs* two. On top of that you pay a launch and teardown per pass, and you materialise two temporaries that exist only to be read once by the next stage. ## What fusion removes A fused pass loads a tile of the tensor into on-chip storage, runs the whole chain on that tile while it sits there, and writes the final result out: ``` fused: read B (+ the residual), write B ``` So what disappears is: the two intermediate writes, the two matching re-reads, the extra launches, and the memory that the temporaries occupied. Because the chain was bandwidth-bound, removing two thirds of its traffic removes roughly two thirds of its time. Nothing about the arithmetic changes — the same normalization, the same activation, the same addition, computed in the same order and to the same values. ## What fusion does not do - **It does not reduce FLOPs.** A candidate who claims fusion "does less maths" has the mechanism wrong. - **It does not help compute-bound work.** Two large dense matmuls back to back spend their time in the arithmetic units; the intermediate they exchange is small relative to that arithmetic, and forcing them into one pass constrains the tiling that keeps those units fed. Fusion is a lever for the left side of the ridge point only. - **It cannot beat the chain's share of the step.** If the chain is 20% of step time and fusion makes it three times faster, step time drops by about 13%. Estimate that ceiling before spending a week on it. - **It does not fix a device that is idle.** If the accelerator is waiting on data that has not arrived, no amount of fusion changes the step time. ## What limits how much you can fuse The intermediates must live in registers or on-chip memory for the tile being processed, so the number of live values per element bounds the chain length. **Reductions are the other constraint**: a normalization's mean and variance depend on every element of a feature vector, so all of that vector's elements must be visited before any of them can be scaled. That forces either a two-pass structure over a tile held on-chip or a single-pass accumulation of running statistics — either is fine as long as the vector fits, but it does mean a reduction is a synchronisation point in a way that purely elementwise stages are not. Chains of elementwise stages, by contrast, compose freely. Finally, the **backward pass** is a separate chain with its own traffic: gradients flow back through the addition, the activation and the normalization, each of which reads and writes tensor-sized data. Fusing only the forward fixes at most half of the problem, and the backward's fused form needs the forward intermediates it depends on, whether they were kept or regenerated. ## How to reason about it in an interview The strong answer sequence is: classify (these are bandwidth-bound, here is the FLOPs-per-byte argument), quantify (the chain moves roughly six tensor-sized transfers, the minimum is two, so the ceiling is about a three-times improvement on this chain), scope (the chain is x% of step time, so the step improves by roughly `x - x/3`), and then state the limits — no help for the matmuls, the reduction constrains the fusion boundary, and the backward needs the same treatment. Recommending a device with a higher FLOP rating for this symptom is the answer that fails the question.
- Why does fusing two large back-to-back matrix multiplications rarely pay off?Because they are compute-bound. Each spends its time in the arithmetic units, and the intermediate they exchange is tiny next to the FLOPs performed, so the traffic you would remove is a small share of the cost. Worse, fusing them constrains the blocking and tile shapes that keep the matrix units saturated, so the fused version often runs slower. Fusion is a lever for operations whose time is dominated by moving bytes.
- How would you estimate the ceiling on a fusion before committing engineering time to it?Compute the bytes the chain moves today and the minimum it could move — one read of the inputs and one write of the output. Divide the difference by achievable bandwidth to get the time you could save, then compare that against the measured step time. If the chain is a small share of the step, the fused version cannot rescue it, and the honest answer is to look elsewhere first.
- What decides how long a fused chain can be?How much intermediate state has to stay resident on-chip per tile. Purely elementwise stages hold one value per element and chain almost indefinitely. A reduction across a dimension — a mean and variance, for instance — requires the whole reduced vector to be visited before its results are used, so the tile must cover that dimension or the pass must accumulate running statistics. Once the live state exceeds on-chip capacity, the intermediates spill and the benefit evaporates.
Unfused is three separate trips to the warehouse to fetch the same crate, do one thing to it, and put it back. Fused is fetching it once, doing all three things at the bench, and returning it.
saying these in an interview costs you the question
- Says fusion reduces the number of floating-point operations
- Expects fusion to speed up a compute-bound matrix multiply
- Recommends a device with a higher FLOP rating for a bandwidth-bound chain
- Forgets the backward pass moves its own tensor-sized traffic
- Claims the whole step gets three times faster, not just the chain