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

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

AnyType, Deinitable, Movable

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)