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