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_level0_band

def fa4_ws_level0_band[band_cols: Int](o_tmem: UInt32, mut o_band: Array[Float32, band_cols])

Level 0 of the shared-key combine: this warp's own O depth band, as-is.

The shared-key counterpart of fa4_ws_level1_combine, and it is startling how little is left. Level 1 exists because the key-split gives each of a warpgroup's m_pack warps an INDEPENDENT key partition carrying a full-depth O partial, so it must stage all m_pack partials through SMEM, barrier, agree a cross-warp max, and fold them with per-partition scales.

Under config.ws_shared_key there is one key band and ONE accumulator, and the packed-TMEM quarters are DEPTH BANDS of the same result. So:

  • No fold. Warp g's band is already final -- there is no O_p from a sibling to add. The m_pack-way FMA chain, the scale[]/lps[] arrays and the whole staging round trip disappear.
  • No barrier, and none is missing. The m_pack warps read DISJOINT hardware subpartitions of one accumulator that a single MMA wrote. There is no cross-warp data flow here at all, so there is nothing to order. (The cross-warp traffic the mode does need is elsewhere: the per-KV-tile row-max and the once-per-tile row-sum, both via fa4_ws_exchange4.)
  • No max. Every warp already holds the same row_max, because the per-iteration fa4_ws_exchange4["max"] made it so. That agreement is the premise the deferred row-sum rests on -- see fa4_ws_exchange4's bit-identity paragraph.

Packed-TMEM addressing rule, stated here rather than cited: for a tcgen05.mma.ws accumulator with MMA_M < 128 and m_pack = 128 // MMA_M, an MMA_M x MMA_N accumulator occupies only MMA_N / m_pack PHYSICAL TMEM columns, because the hardware subpartition does the lane folding. So all m_pack warps issue this tcgen05_ld at the SAME column address and the routing hands each its own column-quarter. Do NOT add a per-warp column offset; the per-warp logical depth base (depth_base + partition_g * band_cols) is a consumer-side quantity and belongs to Level 2, which takes it as depth_base/band_cols.

Caller contract, identical to fa4_ws_intracta_combine: make O visible in TMEM (wait the O-producer barrier and issue tcgen05_fence_after()) before calling. No tcgen05_load_wait() here -- the loads are ordered by data dependency, and the wait belongs to whoever next OVERWRITES the TMEM.

Parameters:

  • โ€‹band_cols (Int): Physical columns of one accumulator this warp reads, i.e. config.pv_mma_n() // config.m_pack. Call once per output depth tile with o_tmem advanced by band_cols; do NOT pass the whole o_phys_cols() at num_o_tiles() > 1, or the register array grows past what num_reg_softmax affords (128 f32 at d512, against 192).

Args:

  • โ€‹o_tmem (UInt32): TMEM column address of THIS depth tile's accumulator, i.e. tmem_addr + config.TMEM_O{wg} + t * band_cols.
  • โ€‹o_band (Array[Float32, band_cols]): Out-param receiving the band, unnormalized, in registers.

Was this page helpful?