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 noO_pfrom a sibling to add. Them_pack-way FMA chain, thescale[]/lps[]arrays and the whole staging round trip disappear. - No barrier, and none is missing. The
m_packwarps 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 viafa4_ws_exchange4.) - No max. Every warp already holds the same
row_max, because the per-iterationfa4_ws_exchange4["max"]made it so. That agreement is the premise the deferred row-sum rests on -- seefa4_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 witho_tmemadvanced byband_cols; do NOT pass the wholeo_phys_cols()atnum_o_tiles() > 1, or the register array grows past whatnum_reg_softmaxaffords (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.