For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
fa4_ws_exchange4
def fa4_ws_exchange4[op: StringSpan[ImmStaticOrigin], *, m_pack: Int, rows: Int](partial: Float32, local_row: UInt32, partition_g: UInt32, warp_group_idx: UInt32, exchange_seq: UInt32, xchg_smem: Pointer[Float32, MutAnyOrigin, address_space=AddressSpace.SHARED]) -> Float32
m_pack-way cross-warp reduction within one softmax warpgroup.
The shared-key sub-mode (config.ws_shared_key) gives all m_pack warps
of a warpgroup ONE key band and ONE shared O accumulator, so the warps
must agree on the running row max every KV tile before they may rescale
it. Warp shuffles cannot cross warps, so the reduction transits SMEM.
"Everyone reduces": every thread writes its own partial, then every thread
reads all m_pack of them. No broadcast step, no designated reducer --
which is what makes ONE barrier enough (see below).
Generalized from the depth512 exchange_reduce ancestor, which is 2-way
and spends two barriers: it single-buffers into correction_smem and
does a read-modify-write there, so the upper half must not overwrite the
slot before the lower half has read it. This one has its own
double-buffered region and no write-back, so one barrier is provably
sufficient.
Why one barrier is race-free. Let p = exchange_seq & 1. Exchange i
writes buffer p, barriers, reads p. A warp that runs ahead writes
1 - p at i+1 -- a different buffer, so it cannot disturb a sibling
still reading p. To touch p again it must reach i+2, which requires
passing i+1's barrier, which the slow sibling must also reach; and to
reach it the sibling must already have finished reading p. So no warp is
ever two exchanges ahead. Do not add a second "to be safe" barrier --
it is pure critical-path cost on the per-KV-tile path.
That argument depends on exchange_seq being a monotone counter across
every call in the warpgroup's lifetime, never reset per phase. In
particular the post-loop row-sum "add" must continue the loop's sequence
rather than restart at 0: restarting could hand it the same parity a
lagging warp is still reading from the final loop iteration.
Bit-identical across warps, by construction. The fold is a fixed left
fold over g = 0 .. m_pack-1, the same order on every thread, so all
m_pack warps obtain the same f32 result -- not merely equal-up-to-
rounding. The shared-key design's deferred row-sum rests on this: with an
identical row_max the per-warp correction = exp2(diff) is identical,
so the warp-uniform lazy-rescale vote resolves the same way everywhere and
the four partial sums stay in one scale, making the final combine a plain
add.
Barrier id is warp_group_idx (0 or 1), the established WG-scoped idiom
in this file. Two facts make it safe and neither is enforced by the type
system:
- The correction, MMA and load warps never call
named_barrier. They do execute one full-CTAbarrier()before the warp-role dispatch, hardware barrier 0 and the onlybarrier()in the kernel, so it is phase-separated from every 128-thread use of id 0, never concurrent. (Under pair-CTA / split-K that site iscluster_sync()instead, which does not touch the named-barrier id space.) - The two hard-coded
named_barrier[...](Int32(0))sites sit past awarp_group_idx != 0 -> returnguard, so only WG0 arrives at them.
Parameters:
- op (
StringSpan[ImmStaticOrigin]):"max"or"add". - m_pack (
Int): Warps per warpgroup sharing the reduction (2 or 4). - rows (
Int): Query rows per warp (BM).m_pack * rowsmust beWARPGROUP_SIZE, since the slot map is a bijection from the warpgroup's 128 threads onto 128 Float32 slots.
Args:
- partial (
Float32): This thread's partial value. - local_row (
UInt32): Query rowrthis thread owns, in[0, rows). - partition_g (
UInt32): This thread's warp index within the warpgroup,[0, m_pack). - warp_group_idx (
UInt32): 0 or 1; doubles as the named-barrier id. - exchange_seq (
UInt32): Monotone per-warpgroup exchange counter (see above). - xchg_smem (
Pointer[Float32, MutAnyOrigin, address_space=AddressSpace.SHARED]): Base ofSM100AttentionSMem.ws_exchange_smem(), i.e. all2 * 2 * WARPGROUP_SIZEslots; this function selects its own warpgroup's and parity's window.
Returns:
Float32: The reduction over all m_pack warps' partials for row local_row,
identical in every warp.