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
SM100AttentionSMem
struct SM100AttentionSMem[qkv_dtype: DType, rope_dtype_: Optional[DType], scale_dtype_: Optional[DType], //, config: FA4Config[qkv_dtype, rope_dtype_=rope_dtype_, scale_dtype_=scale_dtype_], *, use_order_barriers: Bool = EnableForcedOrdering]
Shared memory layout manager for SM100 Flash Attention kernels.
Stores a base pointer into dynamic shared memory and provides accessor methods for each region (Q, K, V, correction, mbarriers, tmem address). All byte-offset arithmetic is comptime so the accessors compile down to a single pointer add + bitcast.
Parametersβ
- βqkv_dtype (
DType): Element type of Q/K/V data in shared memory. - βrope_dtype_ (
Optional[DType]): Element type of Q and K rope (unset when there is no rope). - βscale_dtype_ (
Optional[DType]): Element type of the per-token scale used for Q and K (unset when there is no per-token scaling). - βconfig (
FA4Config[qkv_dtype, rope_dtype_=rope_dtype_, scale_dtype_=scale_dtype_]): FA4 configuration (tile sizes, depths, staging counts, etc.). - βuse_order_barriers (
Bool): Whether forced-ordering barriers are allocated.
Fieldsβ
- βbase (
Pointer[UInt8, MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]):
Implemented traitsβ
AnyType,
Copyable,
ImplicitlyCopyable,
ImplicitlyDeletable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime membersβ
correction_byte_offsetβ
comptime correction_byte_offset = (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth))) + Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))
correction_bytesβ
comptime correction_bytes = (Int((mul _resolve_warpgroup_size(), 2)) * size_of[DType.float32]())
correction_offsetβ
comptime correction_offset = SIMD((Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) // size_of[DType.float32]()))
k_scale_byte_offsetβ
comptime k_scale_byte_offset = (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul _resolve_warpgroup_size(), size_of[DType.float32](), 2), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) + Int((mul size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BM)))
k_scale_bytesβ
comptime k_scale_bytes = (config.num_k_scale_bufs() * Int((mul size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BN)))
k_scale_stride_bytesβ
comptime k_scale_stride_bytes = (config * size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0))
k_stage_bytesβ
comptime k_stage_bytes = SM100AttentionSMem[config, use_order_barriers=use_order_barriers].ws_subtile_bytes if config.use_ws else (Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), config.BN)) * size_of[qkv_dtype]()) if config.use_shared_kv else (Int((mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN)) + Int((mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0))))
k_total_bytesβ
comptime k_total_bytes = (config * Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))))
kv_byte_offsetβ
comptime kv_byte_offset = SM100AttentionSMem[config, use_order_barriers=use_order_barriers].q_bytes
kv_bytesβ
comptime kv_bytes = (Int((mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages)) + Int((mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN))) if config.use_shared_kv else SM100AttentionSMem[config, use_order_barriers=use_order_barriers].kv_stages_bytes
kv_offsetβ
comptime kv_offset = SIMD((Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth))) // size_of[qkv_dtype]()))
kv_stages_bytesβ
comptime kv_stages_bytes = (config * Int((add Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)))))
mbar_byte_offsetβ
comptime mbar_byte_offset = (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul _resolve_warpgroup_size(), size_of[DType.float32](), 2), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), (mul size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BM), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) + Int((mul config.num_k_scale_bufs(), size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BN)))
mbar_bytesβ
comptime mbar_bytes = (SIMD(FA4MiscMBars.num_mbars()) * size_of[SharedMemBarrier]())
mbar_offsetβ
comptime mbar_offset = SIMD((Int((add (mul config.num_k_scale_bufs(), size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BN), (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul _resolve_warpgroup_size(), size_of[DType.float32](), 2), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), (mul size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BM), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) // size_of[SharedMemBarrier]()))
MiscMBarsTypeβ
comptime MiscMBarsType = FA4MiscMBars[num_qk_stages=config.num_qk_stages, num_pv_stages=config.num_pv_stages, num_kv_stages=config.num_kv_stages, use_order_barriers=use_order_barriers, use_shared_kv=config.use_shared_kv, pair_cta=config.pair_cta, num_q=config.num_q, splitk_partitions=config.splitk_partitions, BM=config.BM, use_ws=config.use_ws]
num_k_scale_bufsβ
comptime num_k_scale_bufs = config.num_k_scale_bufs()
num_rope_bufsβ
comptime num_rope_bufs = config.num_rope_buffers()
q_byte_offsetβ
comptime q_byte_offset = Int(0)
q_bytesβ
comptime q_bytes = (Int((mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth)) + Int((mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM)))
q_nope_bytesβ
comptime q_nope_bytes = (Int((mul config.BM, config.padded_nope_depth)) * size_of[qkv_dtype]())
q_offsetβ
comptime q_offset = Int32(0)
q_rope_byte_offsetβ
comptime q_rope_byte_offset = SM100AttentionSMem[config, use_order_barriers=use_order_barriers].q_nope_bytes
q_rope_bytesβ
comptime q_rope_bytes = (Int((mul config.rope_depth(), config.BM)) * size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0))
q_scale_byte_offsetβ
comptime q_scale_byte_offset = (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) + Int((mul _resolve_warpgroup_size(), size_of[DType.float32](), 2)))
q_scale_bytesβ
comptime q_scale_bytes = (config * size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0))
rope_byte_offsetβ
comptime rope_byte_offset = (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth))) + Int((mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages)))
rope_bytesβ
comptime rope_bytes = (Int((mul config.num_rope_buffers(), config.rope_depth(), config.BN)) * size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0))
rope_depthβ
comptime rope_depth = config.rope_depth()
rope_dt_sizeβ
comptime rope_dt_size = size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)
rope_dtypeβ
comptime rope_dtype = rope_dtype_.or_else(qkv_dtype)
rope_stage_elemsβ
comptime rope_stage_elems = (config * config.rope_depth())
scale_dtypeβ
comptime scale_dtype = scale_dtype_.or_else(qkv_dtype)
tmem_addr_byte_offsetβ
comptime tmem_addr_byte_offset = (Int((add (mul config.num_k_scale_bufs(), size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BN), (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul _resolve_warpgroup_size(), size_of[DType.float32](), 2), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth), (mul size_of[(load_from_mem scale_dtype_.value())]() if scale_dtype_.__bool__() else Int(0), config.BM), Int((add (mul config.num_rope_buffers(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BN), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages))) if config.use_shared_kv else Int((add (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages), (mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul config.v_cols_per_cta(), size_of[qkv_dtype](), config.BN)), config.num_kv_stages))))) + Int((mul SIMD(FA4MiscMBars.num_mbars()), size_of[SharedMemBarrier]())))
v_byte_offsetβ
comptime v_byte_offset = SM100AttentionSMem[config, use_order_barriers=use_order_barriers].kv_byte_offset if config.use_shared_kv else (Int((add (mul config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0), config.BM), (mul size_of[qkv_dtype](), config.BM, config.padded_nope_depth))) + Int((mul Int((mul (config.shared_kv_cols() // config), size_of[qkv_dtype](), config.BN)) if config.use_ws else Int((mul max(config.nope_cols_per_cta(), config.v_cols_per_cta()), size_of[qkv_dtype](), config.BN)) if config.use_shared_kv else Int((add (mul config.k_rows_per_cta(), config.rope_depth(), size_of[(load_from_mem rope_dtype_.value())]() if rope_dtype_.__bool__() else Int(0)), (mul config.nope_cols_per_cta(), size_of[qkv_dtype](), config.BN))), config.num_kv_stages)))
v_stage_bytesβ
comptime v_stage_bytes = SM100AttentionSMem[config, use_order_barriers=use_order_barriers].ws_subtile_bytes if config.use_ws else (Int((mul config.v_cols_per_cta(), config.BN)) * size_of[qkv_dtype]())
ws_subtile_bytesβ
comptime ws_subtile_bytes = (Int((mul (config.shared_kv_cols() // config), config.BN)) * size_of[qkv_dtype]())
Methodsβ
__init__β
def __init__() -> Self
Obtain the base pointer from the kernel's dynamic shared memory.
misc_mbarsβ
def misc_mbars(self) -> Self.MiscMBarsType
Return the FA4MiscMBars wrapper over the mbarrier region.
Returns:
Self.MiscMBarsType
q_smemβ
def q_smem(self) -> Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the Q region (offset 0).
Returns:
Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
q_rope_smemβ
def q_rope_smem(self) -> Pointer[Scalar[Self.rope_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the Q rope region (after Q nope in smem).
Returns:
Pointer[Scalar[Self.rope_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
o_smemβ
def o_smem[output_type: DType](self) -> Pointer[Scalar[output_type], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Same physical memory as Q, bitcast to the output element type.
Parameters:
- βoutput_type (
DType): Element type to reinterpret the Q region as.
Returns:
Pointer[Scalar[output_type], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
k_smem_baseβ
def k_smem_base(self) -> Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the K region (first stage, offset = kv_byte_offset).
Returns:
Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
v_smem_baseβ
def v_smem_base(self) -> Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the V region (stage 0).
Non-shared mode: V stage 0 starts after all K stages at kv_byte_offset + num_kv_stages * padded_qk_depth * BN * sizeof. Shared mode: Returns the same pointer as k_smem_base() since K_nope and V share the same buffer.
Returns:
Pointer[Scalar[qkv_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
rope_smem_baseβ
def rope_smem_base(self) -> Pointer[Scalar[Self.rope_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the rope region (shared mode only).
Returns:
Pointer[Scalar[Self.rope_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
correction_smemβ
def correction_smem(self) -> Pointer[Float32, MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the correction region (BM Float32 elements).
Returns:
Pointer[Float32, MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
q_scale_smemβ
def q_scale_smem(self) -> Pointer[Scalar[Self.scale_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the q_scale region (BM elements).
Returns:
Pointer[Scalar[Self.scale_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
k_scale_smemβ
def k_scale_smem(self) -> Pointer[Scalar[Self.scale_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Base of the k_scale region (num_k_scale_bufs * BN elements).
Returns:
Pointer[Scalar[Self.scale_dtype], MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
tmem_addr_ptrβ
def tmem_addr_ptr(self) -> Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
Pointer to the single UInt32 storing the TMEM address.
Returns:
Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.SHARED, _safe=False]
smem_sizeβ
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!