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 function
q_smem_shape
def q_smem_shape[dtype: DType, swizzle_mode: TensorMapSwizzle, *, BM: Int, group: Int, depth: Int, decoding: Bool, fuse_gqa: Bool = False, num_qk_stages: Int = Int(1)]() -> IndexList[Int(4) if decoding or fuse_gqa else Int(3)]
Computes the shared-memory shape for a Q tensor TMA tile based on the tile configuration.
Parameters:
- βdtype (
DType): Element type of the Q tensor. - βswizzle_mode (
TensorMapSwizzle): TMA swizzle mode for the Q tensor tile. - βBM (
Int): Tile block size in the query (row) dimension, in elements. - βgroup (
Int): Grouped-query attention group size, in query heads per KV head. - βdepth (
Int): Head dimension of the attention layer, in elements. - βdecoding (
Bool): Whether the kernel runs in single-token decoding mode. - βfuse_gqa (
Bool): Whether to fuse grouped-query attention into the tile shape (defaults toFalse). - βnum_qk_stages (
Int): Number of pipeline stages used to split the Q shared-memory tile along the depth dimension (defaults to 1).
Returns:
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!