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