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
MLASparseSharedMemoryQKVFP8
struct MLASparseSharedMemoryQKVFP8[config: MLASparseConfig[config.qkv_dtype, config.b_topk_, config.num_mbars_, config.q_smem_depth_, config.q_tmem_depth_, config.cta_group_]]
Native-FP8 SMEM layout: FP8 Q/KV/P + FP8-agnostic softmax scratch.
Unlike MLASparseSharedMemoryFP8, there is no BF16 K/V region and no FP8
staging (the operand IS the FP8 data the MMA reads), so K/V halve in bytes
and the LOCAL dequant staging arrays disappear entirely.
Fieldsβ
- βq (
Array[Float8_e4m3fn, Int((mul (config // cta_group_), config.qk_depth))]): - βkv (
Array[Float8_e4m3fn, (MLASparseSharedMemoryQKVFP8[config].num_mbars * Int((mul (b_topk_ // cta_group_), config.qk_depth)))]): - βp (
Array[Float8_e4m3fn, (MLASparseSharedMemoryQKVFP8[config].num_mbars * Int((mul (config // cta_group_), b_topk_)))]): - βv (
Array[Float8_e4m3fn, (MLASparseSharedMemoryQKVFP8[config].num_mbars * Int((mul (config // cta_group_), b_topk_)) if (eq cta_group_, 2) else Int(128))]): - βd_indices (
Array[Int32, (MLASparseSharedMemoryQKVFP8[config].num_mbars * MLASparseSharedMemoryQKVFP8[config].B_TOPK)]): - βd_indices_v (
Array[Int32, (MLASparseSharedMemoryQKVFP8[config].num_mbars * MLASparseSharedMemoryQKVFP8[config].B_TOPK if MLASparseSharedMemoryQKVFP8[config] else Int(1))]): - βrowwise_max (
Array[Float32, _resolve_warpgroup_size()]): - βrowwise_sum (
Array[Float32, _resolve_warpgroup_size()]): - βis_k_valid (
Array[UInt8, (MLASparseSharedMemoryQKVFP8[config].num_mbars * MLASparseSharedMemoryQKVFP8[config].MASK_BYTES_PER_BUF)]): - βtmem_addr (
Array[UInt32, Int(1)]): - βprologue_q (
Array[SharedMemBarrier, Int(1)]): - βqk_done (
Array[SharedMemBarrier, Int(4)]): - βsv_done (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].num_mbars]): - βkv_ready (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].num_mbars]): - βp_free (
Array[SharedMemBarrier, Int(4)]): - βso_ready (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].num_mbars]): - βk_valid_ready (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].num_mbars]): - βk_valid_free (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].num_mbars]): - βk_ready (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].CG2_MBARS]): - βv_ready (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].CG2_MBARS]): - βv_tma_done (
Array[SharedMemBarrier, MLASparseSharedMemoryQKVFP8[config].CG2_MBARS]):
Implemented traitsβ
comptime membersβ
B_TOPKβ
comptime B_TOPK = config.B_TOPK
B_TOPK_PER_CTAβ
comptime B_TOPK_PER_CTA = (config.B_TOPK // config.cta_group)
CG2_MBARSβ
comptime CG2_MBARS = MLASparseSharedMemoryQKVFP8[config].num_mbars if MLASparseSharedMemoryQKVFP8[config] else Int(1)
INDICES_PER_LANEβ
comptime INDICES_PER_LANE = 8
is_cg2β
comptime is_cg2 = (config.cta_group == Int(2))
KV_STAGE_SIZEβ
comptime KV_STAGE_SIZE = (MLASparseSharedMemoryQKVFP8[config].B_TOPK_PER_CTA * config)
MASK_BYTES_PER_BUFβ
comptime MASK_BYTES_PER_BUF = (config.B_TOPK // Int(8))
NUM_KV_VALID_LANESβ
comptime NUM_KV_VALID_LANES = MLASparseSharedMemoryQKVFP8[config].MASK_BYTES_PER_BUF
num_mbarsβ
comptime num_mbars = config.num_mbars
NUM_Q_HEADSβ
comptime NUM_Q_HEADS = config.num_q_heads
NUM_S_SLOTSβ
comptime NUM_S_SLOTS = 4
O_SIZEβ
comptime O_SIZE = ((config // cta_group_) * config)
P_STAGE_SIZEβ
comptime P_STAGE_SIZE = ((config // cta_group_) * MLASparseSharedMemoryQKVFP8[config].B_TOPK)
PADDED_HEADSβ
comptime PADDED_HEADS = config.padded_num_q_heads
PADDED_HEADS_PER_CTAβ
comptime PADDED_HEADS_PER_CTA = (config // config.cta_group)
Q_SIZEβ
comptime Q_SIZE = ((config // cta_group_) * config)
qk_depthβ
comptime qk_depth = config.qk_depth
S_SLOT_STRIDEβ
comptime S_SLOT_STRIDE = (config.B_TOPK // Int(2))
S_TMEM_BASEβ
comptime S_TMEM_BASE = 256
v_depthβ
comptime v_depth = config.v_depth
V_SMEM_COLS_PER_CTAβ
comptime V_SMEM_COLS_PER_CTA = (config // config.cta_group)
V_STAGE_SIZEβ
comptime V_STAGE_SIZE = (MLASparseSharedMemoryQKVFP8[config].B_TOPK * (config // cta_group_)) if MLASparseSharedMemoryQKVFP8[config] else Int(128)
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!