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 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: