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 function

tma_tile_o

def tma_tile_o[dtype: DType, //, swizzle_mode: TensorMapSwizzle, *, BM: Int, BK: Int, depth: Int](ctx: DeviceContext, ptr: Pointer[Scalar[dtype]], rows: Int, out res: TMATensorTile[dtype, Int(3), IndexList(Int(1), BM, BK, __list_literal__=NoneType(None)), _default_desc_shape[Int(3), dtype, IndexList(Int(1), BM, BK, __list_literal__=NoneType(None)), swizzle_mode]()])

Creates the MLA decode output TMA descriptor.

The row axis is its own descriptor dim of extent BM, reached through a separate outer dim with the same row stride. This makes the stored row count a per-copy coordinate, built by store_row_coords. Box row r lands on tensor row row + r, and box rows at or past rows_to_store fall outside the extent and are dropped. The descriptor base sits BM rows before ptr, but masked rows are never written, so memory before ptr is never touched.

Parameters:

  • ​dtype (DType): Element type of the output tensor (inferred).
  • ​swizzle_mode (TensorMapSwizzle): TMA swizzle mode applied to the descriptor.
  • ​BM (Int): Row extent of the box.
  • ​BK (Int): Tile width in columns for each TMA copy.
  • ​depth (Int): Column count of the full output tensor.

Args:

  • ​ctx (DeviceContext): Device context used to create the TMA descriptor.
  • ​ptr (Pointer[Scalar[dtype]]): Base pointer of the output tensor in device memory.
  • ​rows (Int): Number of rows in the full output tensor.

Returns:

TMATensorTile[dtype, Int(3), IndexList(Int(1), BM, BK, __list_literal__=NoneType(None)), _default_desc_shape[Int(3), dtype, IndexList(Int(1), BM, BK, __list_literal__=NoneType(None)), swizzle_mode]()]

Was this page helpful?