IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo struct

MLA_SM100_Decode_Config

struct MLA_SM100_Decode_Config

Holds the tile sizes, swizzle modes, and SMEM/TMEM layout for an SM100 MLA decode kernel.

Fields​

  • ​MMA_M (Int):
  • ​MMA_PV_N (Int):
  • ​MMA_QK_N (Int):
  • ​BM (Int):
  • ​BN_PV (Int):
  • ​BN_QK (Int):
  • ​BK_QK (Int):
  • ​q_depth (Int):
  • ​depth (Int):
  • ​padded_depth (Int):
  • ​padded_q_depth (Int):
  • ​rope_depth (Int):
  • ​group (Int):
  • ​num_q_heads (Int):
  • ​num_kv_heads (Int):
  • ​tmem_used (Int):
  • ​num_kv_stages (Int):
  • ​smem_used (Int):
  • ​dtype_size (Int):
  • ​num_threads (Int):
  • ​swizzle_mode (TensorMapSwizzle):
  • ​kv_mma_swizzle_mode (TensorMapSwizzle):
  • ​kv_tma_swizzle_mode (TensorMapSwizzle):
  • ​content_swizzle_mode (TensorMapSwizzle):
  • ​rope_swizzle_mode (TensorMapSwizzle):
  • ​decoding_warp_split_k (Bool):
  • ​out_rows (Int):
  • ​page_size (Int):
  • ​split_page_size (Int):
  • ​scale_block_size (Int):
  • ​scales_per_token (Int):
  • ​scale_smem_per_stage (Int):
  • ​per_token_scale_rope_aware (Bool):
  • ​per_token_scales_per_stage (Int):
  • ​decode_layout_g (Bool):
  • ​BK_PV (Int):
  • ​skip_correction_threshold (Float32):

Implemented traits​

AnyType, ImplicitlyDeletable

comptime members​

cta_group​

comptime cta_group = 1

MAX_TMEM_S_SLOTS​

comptime MAX_TMEM_S_SLOTS = Int(6)

mbar_size​

comptime mbar_size = size_of[DType.int64]()

MMA_K​

comptime MMA_K = 16

sm100_smem_carveout​

comptime sm100_smem_carveout = (GPUInfo.from_family(AcceleratorArchitectureFamily(Int(32), Int(2048), Int(233472), Int(65536), Int(1024)), StringSlice("B200"), Vendor(Int8(2)), StringSlice("cuda"), StringSlice("blackwell"), SIMD(10), StringSlice("sm_100a"), Int(148)) - Int(1024))

sm100_tmem_cols​

comptime sm100_tmem_cols = 512

TMEM_CORR_LI​

comptime TMEM_CORR_LI = Int(449)

TMEM_CORR_SCALE​

comptime TMEM_CORR_SCALE = Int(448)

TMEM_O​

comptime TMEM_O = Int(0)

TMEM_S0​

comptime TMEM_S0 = Int(256)

TMEM_S1​

comptime TMEM_S1 = Int(288)

Methods​

__init__​

def __init__(out self, *, num_q_heads: Int, group: Int, depth: Int, q_depth: Int, dtype_size: Int, kv_type_size: Int, swizzle_mode: TensorMapSwizzle, kv_mma_swizzle_mode: TensorMapSwizzle, page_size: Int, decoding_warp_split_k: Bool, split_page_size: Int = Int(128), scale_block_size: Int = Int(0), native_fp8: Bool = False, per_token_scale_rope_aware: Bool = False, decode_layout_g: Bool = False, bn_qk: Int = Int(0), bk_pv: Int = Int(0), skip_correction_threshold: Float32 = -6)

supported​

def supported(self) -> Bool

Returns:

Bool

Was this page helpful?