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):
  • ​nthreads (Int):
  • ​reg_cap (Int):
  • ​reg_consumer (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​

n_tokens​

comptime n_tokens = (mma_n // num_heads)

s_cols​

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

tmem_cols​

comptime tmem_cols = next_power_of_two((Int(2) * IndexPrefillConfig[dtype, ks_dtype, depth, BM_key, mma_n, num_heads].s_cols))

Methods​

__init__​

def __init__() -> Self

Was this page helpful?