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_ws_ts_partial
def bulk_mma_ws_ts_partial[kind: UMMAKind, b_dtype: DType, *, b_BMN: Int, b_BK: Int, b_swizzle: TensorMapSwizzle, b_is_k_major: Bool, num_k_mmas: Int, operand_size: Int, tcgen05_mma_type: String, mma_k: Int = Int(16), k_start: Int = Int(0), b_page_dense: Bool = False](idesc: UMMAInsDescriptor[kind], a: UInt32, b: MMASmemDescriptorPair, c_tmem: UInt32, c_scale: UInt32, elect: Int32, valid_k_mmas: UInt32)
Issues a partial-K TS warp-specialized contraction for a partially-loaded last KV tile.
a is the un-offset TMEM base; each block's absolute column offset is computed in-PTX, and a %pv validity guard is kept separate from elect.
Parameters:
- βkind (
UMMAKind):UMMAKindselecting thetcgen05.mma.wsinstruction variant. - βb_dtype (
DType): Element dtype of the B operand, used to derive the B SMEM tile layout. - βb_BMN (
Int): M (or N) dimension of the B operand tile in elements, used to derive the B layout. - βb_BK (
Int): K dimension of the B operand tile in elements, used to derive the B layout. - βb_swizzle (
TensorMapSwizzle): SMEM swizzle mode for the B tile. - βb_is_k_major (
Bool): Whether B is stored k-major (True) or mn-major (False). - βnum_k_mmas (
Int): Number ofmma_k-sized K-dimension blocks to contract over in this stage. - βoperand_size (
Int): Size in bytes of the A and B operand elements. - βtcgen05_mma_type (
String):tcgen05.mma.wsinstruction string prefix, including the CTA-group selector. - βmma_k (
Int): K-dimension tile size per MMA block, in elements (defaults to 16). - βk_start (
Int): Absolute K-block index of the first block in this stage (defaults to 0). - βb_page_dense (
Bool): Whether B uses the row-major page-fold layout (defaults toFalse).
Args:
- βidesc (
UMMAInsDescriptor[kind]): UMMA instruction descriptor encoding the accumulator and operand dtypes and the output tile shape. - βa (
UInt32): Un-offset (stage-0) TMEM base address of 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!