How does tensor parallelism split a transformer feed-forward block's two matmuls across devices?
answer
- one matmul widens, the next narrows
- columns first, rows second
- elementwise activation needs no neighbours
- partials summed, not concatenated
- one all-reduce at the block's exit
basics
~20 sSplit the first weight matrix by columns so each device computes its own slice of the hidden activations, then split the second by rows so each device produces a partial sum of the output. One all-reduce adds the partials.
solid answer
~50 sA feed-forward block is two matmuls with an elementwise nonlinearity between them: `h = act(x W1)` widens the hidden size from `d` to `4d`, and `y = h W2` projects back down. Split `W1` by columns, so device `i` holds a `d x (4d/N)` slice and computes only its own columns of `h`. Because the nonlinearity is elementwise, each device applies it to its slice locally — nothing needs to be exchanged in the middle. Then split `W2` by rows to match: device `i` holds `(4d/N) x d` and multiplies its slice of `h`, producing a full-width partial sum of `y`. A single all-reduce sums the partials into the true output. That is one collective in the forward pass and one in the backward pass for the input gradient, and each device stores only `1/N` of the block's weights and optimizer state.
go deeper
Recall that tensor parallelism cuts one weight matrix into pieces across devices instead of cutting the batch, and that the pieces have to be recombined before the block's output can be used.
Be ready to state the shapes: a column split for the widening matmul, a matching row split for the projection, and why an elementwise nonlinearity between them needs no communication.
Show that you know where the collective lands in both passes, what its volume depends on, and why the group has to stay inside a fast interconnect for this to be worth doing.
Own how wide to make the tensor-parallel group. It buys memory per device at the cost of blocking latency in every block, and it stops paying the moment the group leaves the high-bandwidth domain or exceeds the head count.
## The problem Tensor parallelism (also called intra-layer model parallelism) splits a single weight matrix across several devices so that no device has to hold the whole layer. It is the tool you reach for when one layer, on its own, is too large for one device's memory — not when the batch is too large. Splitting the batch is a different strategy entirely. ## The shape of the block A transformer's position-wise feed-forward block is two dense layers with an elementwise nonlinearity between them. Write the token activations as `x` with shape `(T, d)`, where `T` is the number of tokens in the micro-batch and `d` is the model's hidden size. The block computes: - `h = act(x W1)` with `W1` of shape `(d, 4d)` — the widening matmul; `h` has shape `(T, 4d)`. - `y = h W2` with `W2` of shape `(4d, d)` — the projection back down; `y` has shape `(T, d)`. The expansion factor of four is conventional, not required; what matters is that the middle dimension is large and is the one you split. ## Column split, then row split With `N` devices, partition the intermediate dimension `4d` into `N` chunks of size `4d/N`. **First matmul, column-parallel.** Device `i` holds the column slice `W1_i` of shape `(d, 4d/N)`. It needs the full input `x`, which is replicated on every device, and computes `h_i = act(x W1_i)` of shape `(T, 4d/N)`. That is a genuine slice of the true `h` — columns `i*4d/N` through `(i+1)*4d/N`. **The nonlinearity is free.** This is the point of choosing a column split first. An elementwise function — GELU, ReLU, or a gated variant computed from the same slice — depends only on the value at each position, never on its neighbours along the hidden dimension. Each device applies it to its own columns with no communication at all. Had you split the first matmul by rows instead, you would have partial sums rather than a slice, and you would have to reduce them *before* the nonlinearity, adding a collective in the middle of the block. **Second matmul, row-parallel.** Device `i` holds the row slice `W2_i` of shape `(4d/N, d)` — exactly the rows that pair with the columns it already owns. It computes `y_i = h_i W2_i`, which has the *full* output shape `(T, d)` but is only a partial sum: the true `y` equals the sum of the `y_i` over all devices, because a matmul over a split inner dimension decomposes into a sum of products over the pieces. **One collective.** An all-reduce over the tensor-parallel group sums the `y_i` and leaves every device with the identical full `y`, ready for the residual add and the following normalization layer, both of which stay replicated. ## The backward pass The two boundaries are conjugates of each other. At the exit, forward is an all-reduce and backward is the identity: the gradient arriving at `y` is already the same on every device, and each device uses it directly. At the entry, forward is the identity (`x` is broadcast as-is) and backward is an all-reduce: each device computes a partial gradient with respect to `x` from its own column slice, and those partials must be summed before the gradient leaves the block. So a tensor-parallel block costs one all-reduce forward and one backward per step. Weight gradients need no exchange — each device's slice of `W1` and `W2` is owned by that device alone, and it updates its own slice with its own optimizer state. ## What it costs, and when it stops paying The communication volume of each all-reduce scales with `T * d` — tokens times hidden size — so it grows with the micro-batch and sequence length, not with how many parameters the layer has. Crucially, that collective is *blocking* and sits on the critical path of every single block: the next layer cannot begin until it finishes. There is no useful work to overlap it with inside the block. That is why the tensor-parallel group is normally kept inside one high-bandwidth domain — a single machine's fast device-to-device links. Stretch the group over a slower interconnect and the per-block latency, multiplied by every block and every micro-batch, dominates the step. The same pattern applies to multi-head attention: the query, key and value projections are column-parallel so that each device owns a disjoint subset of heads and can run their attention entirely locally, and the output projection is row-parallel, again ending in one all-reduce. Heads are independent, so no head is ever split across devices — the head count bounds how wide the tensor-parallel group can usefully be. ## Common failure The two splits must be chosen as a matched pair. Column-then-row gives one collective per block; row-then-column, or column-then-column, forces a reshuffle in the middle and roughly doubles the communication for no benefit.
- What would go wrong if you split the first matmul by rows instead of by columns?A row split of the widening matmul makes each device produce a partial sum of the full-width hidden tensor rather than a slice of it. The nonlinearity is elementwise but not linear, so you cannot apply it to a partial sum — you would have to all-reduce before the activation and again after the second matmul, doubling the collectives per block and gaining nothing.
- How does the same idea apply to multi-head attention?The query, key and value projections are column-parallel so each device owns a disjoint subset of heads, and it runs those heads' attention end to end locally because heads never interact. The output projection is row-parallel, so the block again ends in one all-reduce. Since no head is split, the number of heads caps how wide the tensor-parallel group can usefully be.
- Do the weight gradients need a collective in a tensor-parallel block?No. Each device owns a disjoint slice of both weight matrices, so it computes the gradient for its own slice and updates it with its own optimizer state. The collectives are only for activations and the input gradient: an all-reduce of the block output in the forward pass, and an all-reduce of the gradient with respect to the block input in the backward pass.
Think of the block as an hourglass that widens then narrows. Each device takes one vertical strip of the widening half and the matching strip of the narrowing half, so the strips only have to meet once, at the very bottom.
saying these in an interview costs you the question
- Describes splitting the batch rather than the weight matrices
- Says the per-device outputs are concatenated instead of summed
- Inserts a collective between the two matmuls
- Claims the nonlinearity needs the full hidden vector
- Thinks weight gradients also need an all-reduce here
- Assumes the collective volume scales with parameter count