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

MLAConfig

struct MLAConfig[qkv_dtype: DType, *, rope_gmem_dtype: DType, rope_mma_dtype: DType, scale_dtype_: Optional[DType] = None]

Fields​

  • ​fa4_config (FA4Config[qkv_dtype, rope_dtype_=rope_mma_dtype, scale_dtype_=scale_dtype_]):
  • ​MMA_M (Int):
  • ​BM (Int):
  • ​BN (Int):
  • ​BK0 (Int):
  • ​BK1 (Int):
  • ​qk_depth (Int):
  • ​rope_depth (Int):
  • ​nope_depth (Int):
  • ​cache_depth (Int):
  • ​padded_qk_depth (Int):
  • ​group (Int):
  • ​num_q_heads (Int):
  • ​num_kv_heads (Int):
  • ​TMEM_S1 (Int):
  • ​TMEM_O0 (Int):
  • ​TMEM_O1 (Int):
  • ​TMEM_P0 (Int):
  • ​TMEM_P1 (Int):
  • ​tmem_used (Int):
  • ​num_kv_stages (Int):
  • ​num_qk_stages (Int):
  • ​num_pv_stages (Int):
  • ​smem_used (Int):
  • ​qkv_swizzle_mode (TensorMapSwizzle):
  • ​rope_mma_swizzle_mode (TensorMapSwizzle):
  • ​rope_gmem_swizzle_mode (TensorMapSwizzle):
  • ​output_swizzle_mode (TensorMapSwizzle):

Implemented traits​

AnyType, Copyable, ImplicitlyCopyable, ImplicitlyDeletable, Movable, RegisterPassable, TrivialRegisterPassable

comptime members​

mbar_size​

comptime mbar_size = size_of[DType.int64]()

num_correction_cols​

comptime num_correction_cols = 1

num_threads​

comptime num_threads = Int(512)

qkv_dtype_size​

comptime qkv_dtype_size = size_of[qkv_dtype]()

rope_gmem_dtype_size​

comptime rope_gmem_dtype_size = size_of[rope_gmem_dtype]()

rope_mma_dtype_size​

comptime rope_mma_dtype_size = size_of[rope_mma_dtype]()

scale_dtype​

comptime scale_dtype = scale_dtype_.or_else(qkv_dtype)

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_S0​

comptime TMEM_S0 = Int(0)

Methods​

__init__​

def __init__(*, num_q_heads: Int, group: Int, depth: Int, page_size: Int, v_depth: Int = Int(-1), num_q: Int = Int(2), single_o: Bool = False, bn_cap: Int = Int(0)) -> Self

num_q​

def num_q(self) -> Int

Returns:

Int

q_tile_rows​

def q_tile_rows(self) -> Int

Rows per Q TMA tile / per-half MMA β€” BM // num_q.

128 in both modes: one of two BM=256 halves in 2Q, the single full BM=128 tile in 1Q. The Q (and per-token q_scale) TMA boxes and the ragged output store all fold to this value, which is why their op types match across the 1Q/2Q configs.

Returns:

Int

with_num_q​

def with_num_q(self, num_q: Int) -> Self

Reconstruct this config with a different num_q (single-CTA).

Mirrors FA4Config.with_num_q, but simpler: MLA pins num_qk_stages == 1 (is_mla), so there is no staging knob to match between the 1Q and 2Q variants.

v_depth (the V/output head dim, carried by fa4_config.ov_depth) MUST be re-passed: otherwise the rebuilt config defaults to v_depth == nope_depth, so the 1Q variant's V/output geometry would diverge from the 2Q config's (a type mismatch in the shared O-store / V TMA tile when v_head_dim != qk_nope_head_dim).

switch_1q_config​

def switch_1q_config(self) -> Self

The 1Q variant used by the in-kernel per-sequence 1Q/2Q switch.

Identical to with_num_q(1) (see with_num_q for why MLA has no staging-pinning concern, unlike FA4Config.switch_1q_config).

can_switch_to_1q​

def can_switch_to_1q(self) -> Bool

Whether a 2Q-launched kernel may dispatch to the 1Q body at runtime.

True only when this is a 2Q config AND the 1Q variant is valid. The TMA-op types fold between the two configs by construction: the Q TMA / ragged-store BM // num_q is 128 in both modes, and the K_nope/K_rope/V TMA shapes are BM-independent (BN's formula does not reference num_q).

Returns:

Bool

launch_smem_used​

def launch_smem_used(self) -> Int

Dynamic smem to reserve when launching this config's kernel.

When the launched kernel may dispatch to the 1Q body at runtime (can_switch_to_1q()), it constructs the 1Q SM100AttentionSMem over the same dynamic smem region, so the launch must reserve the max of both footprints. Otherwise this is just smem_used.

Returns:

Int

launch_num_threads​

def launch_num_threads(self) -> Int

Threads to launch for this config's kernel.

The generic single-O (wide-V) path drops the redundant 2nd softmax warpgroup -- WG1 is a full no-op there (see the single-O serial-KV accumulation) -- so it launches 3 warpgroups (Softmax0 + Correction + MMA/Load/Empty) instead of 4. Every other config keeps the standard 4-warpgroup (num_threads = 512) layout. Only the generic kernel calls this; the per-token-scale / blockscale siblings read the num_threads field directly and stay at 512 even for their own single-O configs (they keep the 2nd softmax WG).

Returns:

Int

prefer_1q​

def prefer_1q(self, max_prompt_len: UInt32, num_partitions: UInt32, batch_size: UInt32, sm_count: Int) -> Bool

Runtime 1Q-vs-2Q grid heuristic for a 2Q config (mirrors the MHA heuristic in dispatch.mojo): prefer 1Q when (a) max_prompt_len fits a single 1Q tile (q_tile_rows()), so 2Q's BM=256 would waste >= 50% of Q rows, or (b) the unclamped 2Q grid only fills <= half the SMs, so halving BM doubles the grid without oversubscribing.

Returns:

Bool

num_rope_buffers​

def num_rope_buffers(self) -> Int

Returns:

Int

supported​

def supported(self) -> Bool

Returns:

Bool

correction_smem_elements​

def correction_smem_elements(self) -> Int

Returns:

Int

Was this page helpful?