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

load_matrix_b_amd

def load_matrix_b_amd[m: Int, n: Int, k: Int](b_ptr: Pointer[Float32, _safe=True], tile_row: Int, tile_col: Int, ldm: Int) -> Float32

Loads a tile of matrix B from memory to registers for AMD FP32 tensor core operations.

Parameters:

  • ​m (Int): Number of rows in the output matrix tile.
  • ​n (Int): Number of columns in the output matrix tile.
  • ​k (Int): Inner dimension for matrix multiplication.

Args:

  • ​b_ptr (Pointer[Float32, _safe=True]): Pointer to matrix B data in memory.
  • ​tile_row (Int): Starting row index of the tile.
  • ​tile_col (Int): Starting column index of the tile.
  • ​ldm (Int): Leading dimension of matrix B (stride between rows).

Returns:

Float32: SIMD vector containing 1 FP32 value loaded from matrix B.

def load_matrix_b_amd[dtype: DType, //, m: Int, n: Int, k: Int, n_blocks: Int = Int(1)](b_ptr: Pointer[Scalar[dtype], _safe=True], tile_row: Int, tile_col: Int, ldm: Int, tile_loops: Int = Int(1)) -> SIMD[dtype, SIMDLength(4)] where dtype.is_half_float()

Loads a tile of matrix B from memory to registers for AMD half-precision tensor core operations.

This function loads 4 consecutive values per thread from matrix B in a pattern optimized for AMD GPU tensor core operations. Each thread loads values based on its position within the warp.

Constraints:

The tile dimensions must be m=16, n=16, k=16 and n_blocks=1 or m=4, n=4, k=4 and n_blocks=16.

Parameters:

  • ​dtype (DType): Data type of the matrix elements (float16 or bfloat16).
  • ​m (Int): Number of rows in the output matrix tile.
  • ​n (Int): Number of columns in the output matrix tile.
  • ​k (Int): Inner dimension for matrix multiplication.
  • ​n_blocks (Int): Number of blocks.

Args:

  • ​b_ptr (Pointer[Scalar[dtype], _safe=True]): Pointer to matrix B data in memory.
  • ​tile_row (Int): Starting row index of the tile.
  • ​tile_col (Int): Starting column index of the tile.
  • ​ldm (Int): Leading dimension of matrix B (stride between rows).
  • ​tile_loops (Int): Number of tile loops across matrix B's row dimension.

Returns:

SIMD[dtype, SIMDLength(4)]: SIMD vector containing 4 values loaded from matrix B.