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

Mojo struct

IndexPrefillConfig

struct IndexPrefillConfig[dtype: DType, ks_dtype: DType, depth: Int, BM_key: Int, mma_n: Int, num_heads: Int]

Every derived quantity of one prefill instantiation, in one object.

This is the FA4Config pattern from nvidia/sm100/attention.mojo, adopted for the same reason it exists there: the pipeline's shape, its SMEM accounting and its register accounting are ONE derivation, and splitting them across free helpers is what lets two of them disagree.

What the split cost here, concretely: the fixed-SMEM term fed only the ring-depth derivation while the launcher passed a separate total, so the two had to be edited together by hand -- and a one-sided edit under-allocates the launch rather than failing it. Both are now the same accumulation: smem_used == fixed_smem_bytes + k_stages * k_stage_bytes identically.

Residency stays OUTSIDE the struct, in _ctas_per_sm, because the router asks what residency a tile would take before it has a dtype or depth to build a config with.

Parameters​

  • ​dtype (DType): Q/K element type.
  • ​ks_dtype (DType): K-scale element type.
  • ​depth (Int): Head dimension.
  • ​BM_key (Int): Keys per K tile.
  • ​mma_n (Int): The N tile, n_tokens * num_heads.
  • ​num_heads (Int): Index heads.

Fields​

  • ​ctas_per_sm (Int):
  • ​cons_wgs (Int):
  • ​prod_warps (Int):
  • ​nthreads (Int):
  • ​reg_cap (Int):
  • ​reg_consumer (Int):
  • ​reg_producer (Int):
  • ​static_reg_budget (Int):
  • ​hoist_q_scales (Bool):
  • ​k_stages (Int):
  • ​n_mbars (Int):
  • ​carveout (Int):
  • ​fixed_smem_bytes (Int):
  • ​k_stage_bytes (Int):
  • ​smem_used (Int):

Implemented traits​

AnyType, Copyable, Deinitable, ImplicitlyCopyable, Movable, RegisterPassable, TrivialRegisterPassable

comptime members​

cta_token_stride​

comptime cta_token_stride = (IndexPrefillConfig[dtype, ks_dtype, depth, BM_key, mma_n, num_heads].n_tokens * _s_rings[mma_n]())

fixed_mbars​

comptime fixed_mbars = (Int((mul _s_rings[mma_n](), 4)) + _s_rings[mma_n]() if _force_narrow[mma_n]() else Int(0))

mma_warps​

comptime mma_warps = _mma_warps[mma_n]()

n_tokens​

comptime n_tokens = (mma_n // num_heads)

q_copies​

comptime q_copies = IndexPrefillConfig[dtype, ks_dtype, depth, BM_key, mma_n, num_heads].s_rings

s_cols​

comptime s_cols = align_up(mma_n, Int(32))

s_rings​

comptime s_rings = _s_rings[mma_n]()

split_ks​

comptime split_ks = _split_ks[mma_n]()

stage_mbars​

comptime stage_mbars = Int(4) if IndexPrefillConfig[dtype, ks_dtype, depth, BM_key, mma_n, num_heads].split_ks else Int(2)

tmem_cols​

comptime tmem_cols = next_power_of_two(Int((mul align_up(mma_n, Int(32)), _s_rings[mma_n](), 2)))

two_q​

comptime two_q = _force_narrow[mma_n]()

Methods​

__init__​

def __init__() -> Self

Was this page helpful?