IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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-CTA barrier() before the warp-role dispatch, hardware barrier 0 and the only barrier() 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 is cluster_sync() instead, which does not touch the named-barrier id space.)
  • The two hard-coded named_barrier[...](Int32(0)) sites sit past a warp_group_idx != 0 -> return guard, 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 * rows must be WARPGROUP_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 row r this 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 of SM100AttentionSMem.ws_exchange_smem(), i.e. all 2 * 2 * WARPGROUP_SIZE slots; 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.

Was this page helpful?