For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo function
bulk_mma_ss_partial
def bulk_mma_ss_partial[kind: UMMAKind, //, layout_a: Layout, layout_b: Layout, *, num_k_mmas: Int, mma_k: Int, operand_size: Int, k_start: Int = Int(0), cta_group: Int = Int(1)](idesc: UMMAInsDescriptor[kind], a: MMASmemDescriptorPair, b: MMASmemDescriptorPair, c_tmem: UInt32, c_scale: UInt32, elect: Int32, valid_k_mmas: UInt32)
Issues a partial-K SS contraction for a partially-loaded last KV tile, non-warp-specialized.
Both A and B come from SMEM descriptors; each block's MMA carries a warp-uniform validity guard derived from valid_k_mmas, kept separate from elect.
Parameters:
- βkind (
UMMAKind):UMMAKindselecting thetcgen05.mmainstruction variant. - βlayout_a (
Layout): SMEM layout of the A operand tile, used to compute per-K-block A descriptor offsets. - βlayout_b (
Layout): SMEM layout of the B operand tile, used to compute per-K-block B descriptor offsets. - βnum_k_mmas (
Int): Number ofmma_k-sized K-dimension blocks to contract over in this stage. - βmma_k (
Int): K-dimension tile size per MMA block, in elements. - βoperand_size (
Int): Size in bytes of the A and B operand elements. - βk_start (
Int): Absolute K-block index of the first block in this stage (defaults to 0). - βcta_group (
Int): Number of cooperating CTAs, 1 or 2 (defaults to 1).
Args:
- βidesc (
UMMAInsDescriptor[kind]): UMMA instruction descriptor encoding the accumulator and operand dtypes and the output tile shape. - βa (
MMASmemDescriptorPair): Un-offset (stage-0) SMEM descriptor pair for the A operand. - βb (
MMASmemDescriptorPair): Un-offset (stage-0) SMEM descriptor pair for the B operand. - βc_tmem (
UInt32): TMEM base address of the output accumulatorC. - βc_scale (
UInt32): Accumulator init/accumulate scale; nonzero on the first block to initialize the accumulator, zero to accumulate. - βelect (
Int32):elect()result selecting the single thread that issues the MMA. - βvalid_k_mmas (
UInt32): Count of loadedmma_k-sized blocks; blocks whose absolute index reaches or exceeds this count are predicated off.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!