skip to content

Factorized Layer Design

Splitting an expensive layer into cheaper pieces: a low-rank product in place of a dense matrix, expand-then-project bottlenecks. Interviewers ask where the saving actually comes from.

on this pageshow

questions

4

How does factorizing a 4096x4096 dense layer into two rank-256 matrices cut parameters?

level: middleimportance: must knowfreq 50%

answer

  1. one wide matrix becomes two thin ones
  2. count weights before and after
  3. break-even rank exists
  4. half the width, square layer
  5. the product is capped at r

basics

~20 s

One 4096x4096 matrix holds about 16.8M weights; replacing it with a 4096x256 matrix times a 256x4096 matrix holds about 2.1M, an 8x cut. The saving exists only while the rank stays below half the layer width.

solid answer

~50 s

A dense layer computes `y = W x + b` with `W` of shape m-by-n, costing `m * n` weights and `m * n` multiply-accumulates per input. Factorizing means writing `W` as a product `A B` with `A` m-by-r and `B` r-by-n, so the layer becomes two chained matrix multiplies through an r-wide waist. Cost drops from `m * n` to `r * (m + n)` in both weights and multiply-accumulates. For m = n = 4096 and r = 256 that is 16,777,216 down to 2,097,152, roughly 8x. The break-even rank is `r = m * n / (m + n)`, which for a square layer is half the width, so anything at or above 2048 here saves nothing. The price is expressivity: the layer can now only realise linear maps of rank at most r, and below some rank floor the waist becomes an information bottleneck the rest of the network cannot compensate for.

go deeper

for a junior

Be able to say what the two matrices are and roughly why the count drops: one big rectangle of numbers becomes two thin ones sharing a narrow waist. Recall the 4096-to-rank-256 example and the 8x figure.

for a middle

Derive r * (m + n) versus m * n on the spot and state the break-even rank as m * n / (m + n). Explain that multiply-accumulates fall in the same ratio and that the product is capped at rank r.

for a senior

Show you would count parameters per layer first, target only the widest, and always budget a recovery fine-tune. Expect to discuss where the rank floor sits and how you would find it empirically rather than by rule of thumb.

for a principal

Own the framing that factorization buys size and arithmetic with expressivity, and that the exchange rate is layer-specific. Be ready to argue when the compression target is better met by changing the architecture than by squeezing an existing one.

## The shape of the trick A fully connected layer computes `y = W x + b`. If the input has n units and the output m units, `W` is an m-by-n matrix: `m * n` learned numbers, and `m * n` multiply-accumulate operations for every input vector pushed through it. In wide networks this single matrix is often the dominant term in the model's size. Factorizing the layer means replacing that one matrix with a product of two thinner ones: `W ~= A B`, where `A` is m-by-r and `B` is r-by-n, and r (the rank, or waist width) is much smaller than either dimension. At inference you never form `W` again; you compute `B x` first, getting an r-vector, then `A (B x)`. The layer is now two chained linear maps with a narrow middle, and no nonlinearity between them. ## The arithmetic Parameters go from `m * n` to `r * m + r * n = r * (m + n)`. Multiply-accumulates per input follow exactly the same ratio, because each matrix multiply costs one MAC per weight. For the canonical case, a 4096-by-4096 layer: - dense: 4096 * 4096 = 16,777,216 weights - rank 256: 256 * (4096 + 4096) = 2,097,152 weights - ratio: 8.0x smaller The compression factor is `m * n / (r * (m + n))`. For a square layer of width d that simplifies to `d / (2 r)`: rank 256 on a 4096-wide layer gives 4096 / 512 = 8. ## Break-even rank Factorization only helps while `r * (m + n) < m * n`, i.e. while `r < m * n / (m + n)` For the square 4096 layer that threshold is 2048, exactly half the width. Choose rank 3000 and you have made the layer *larger* while also making it less expressive, which is the worst of both. The threshold is the harmonic-mean-flavoured quantity `m*n/(m+n)`, so rectangular layers behave differently: a 4096-by-1024 layer breaks even at 4096*1024/5120 = 819, and a very lopsided layer (say 50,000-by-768) breaks even near 756 - close to the smaller dimension, which is why tall, thin matrices such as vocabulary tables are the easiest wins. ## What you give up A product `A B` has rank at most r. Whatever the original layer was doing, the factorized version can only produce linear maps whose output lives in an at-most-r-dimensional subspace of the output space. Two consequences follow. First, if the trained map genuinely needed more than r independent directions, that information is simply gone - and it is gone for every downstream layer, so the error does not stay local. Second, there is a rank floor: compression degrades gracefully for a while and then falls off a cliff, because the waist stops being a mild restriction and becomes a hard information bottleneck through which the whole activation must pass. Where that floor sits is layer-specific and must be measured, not guessed. Note also that stacking two matrices with nothing between them does not make the network deeper in any expressive sense. `A B` is still one linear map. The only things that change are the parameter count, the arithmetic cost, the rank ceiling, and the optimization dynamics. ## Where the technique pays The saving is proportional to the layer's size, so the first thing to do is count parameters per layer and attack the widest few. Typical candidates are large hidden-to-hidden projections, big classifier heads, and input embedding tables. In a large-vocabulary model, a V-by-768 embedding table can be factorized into a V-by-128 table times a 128-by-768 projection; for V = 50,000 that is 38.4M weights down to about 6.5M, and it decouples the width of the token vectors from the model's hidden width. Narrow layers rarely repay the effort: with a small `m*n/(m+n)` threshold there is almost no room between a usable rank and break-even. ## Practicalities Bias vectors are untouched - they live on the output side and cost m numbers either way. Any layer normalization or activation that followed the dense layer still follows the projection; the factorization sits strictly inside the linear part. And a factorized layer taken from a trained model should always be followed by a recovery fine-tune: even a good approximation shifts every downstream activation distribution slightly, and a short retrain claws most of that back.

  • At what rank does factorizing an m-by-n weight matrix stop saving anything?
    At `r = m * n / (m + n)`. Below it, `r * (m + n) < m * n` and you save; at or above it the two factors hold more weights than the original. A square layer of width d breaks even at d/2, so a 4096-wide layer must stay under 2048. A 4096-by-1024 layer breaks even at 819, and a very tall, thin matrix breaks even just under its smaller dimension.
  • How does the same trick shrink a large-vocabulary input embedding table?
    Split the V-by-H table into a V-by-E table and an E-by-H projection, with E much smaller than H. Each token id looks up an E-dimensional row, which is then projected up to the model width. For V = 50,000 and H = 768, choosing E = 128 takes 38.4M weights to about 6.5M, and the per-token width no longer has to equal the hidden width.
  • Why doesn't inserting the second matrix make the layer deeper in any useful sense?
    Because there is no nonlinearity between the factors. The composition `A B` is itself a single linear map, just one restricted to rank r, so the layer's function class shrinks rather than grows. What does change is the parameterization: gradients now flow through a product, and the optimization behaviour of a factorized layer is not the same as that of the dense layer it approximates, even though the function family is a subset of it.

Routing every message in a company through a small mailroom: if the mailroom has enough desks, everything still gets where it is going for far less overhead - but shrink it past a point and traffic that used to flow in parallel is permanently squeezed out.

saying these in an interview costs you the question

  • Says two smaller matrices are always cheaper than one big one
  • Picks a rank above half the layer width and calls it compression
  • Claims the factorized layer can still represent the original map
  • Thinks the inserted matrix adds depth and therefore capacity
  • Assumes any rank is fine as long as you fine-tune afterwards
  • Factorizes small layers where almost none of the parameters live

context

open as a page

Why does an inverted residual block expand the channel count with a 1x1 before projecting back?

level: middleimportance: should knowfreq 42%

basics

~20 s

Because the block's middle operator is cheap per channel, so extra width there costs little while giving the nonlinearity room to work. The expensive parts and the residual stay on the narrow ends, which keeps parameters and stored activations small.

open as a page

How does a trained weight matrix's singular-value spectrum tell you whether to factorize it?

level: seniorimportance: should knowfreq 35%

basics

~20 s

By how fast the singular values decay. A steep decay means a few directions carry the map, so a low-rank replacement loses little; a flat, well-conditioned spectrum means every direction matters and factorizing at any useful rank will cost accuracy.

open as a page

Do you factorize a trained model and retrain, or train the factorized shape from scratch?

level: principalimportance: nice to knowfreq 28%

basics

~20 s

Factorize-then-fine-tune when trained weights exist and compute is short: the spectrum picks a per-layer rank and a brief retrain recovers most accuracy. Train the factorized shape from scratch when you want low rank learned, not imposed.

open as a page