Sharded Matmul Simulator
Configure A[I, J] · B[J, K] → C[I, K] and watch the mesh compute it: which collectives run, how shards move, and where the time goes.
Presets:
Case 3 reduction
Case 3: both inputs sharded along the contracting dimension
Each device can multiply its matching J-chunks, but the result is only a partial sum, written C {U_axis}. An AllReduce (or a cheaper ReduceScatter, if a sharded output is fine) completes the sum.
A[I, J_X] · B[J_X, K] → C[I, K]
1 / 4
layoutInitial layout: A[I, J_X] · B[J_X, K] on a 4x4 mesh.
Y=0Y=1Y=2Y=3
X=0X=1X=2X=3
TPU 0 (0,0)
?
CTPU 1 (0,1)
?
CTPU 2 (0,2)
?
CTPU 3 (0,3)
?
CTPU 4 (1,0)
?
CTPU 5 (1,1)
?
CTPU 6 (1,2)
?
CTPU 7 (1,3)
?
CTPU 8 (2,0)
?
CTPU 9 (2,1)
?
CTPU 10 (2,2)
?
CTPU 11 (2,3)
?
CTPU 12 (3,0)
?
CTPU 13 (3,1)
?
CTPU 14 (3,2)
?
CTPU 15 (3,3)
?
CBlocks are coloured by data identity; striped blocks are partial sums awaiting a reduction. Arriving blocks fly in from the device that sent them.
Cost model (TPU v5e, bf16)
| Operation | Bytes | Time |
|---|---|---|
| AllReduceX (C) | 8.39 MB | 186 µs |
| Local matmul / device | 4.29 GFLOPs | 21.8 µs |
Comms total: 186 µsCompute total: 21.8 µsCommunication bound