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
batched_matmul_dynamic_scaled_fp8_naive
def batched_matmul_dynamic_scaled_fp8_naive[c_type: DType, a_type: DType, b_type: DType, a_scales_type: DType, b_scales_type: DType, //, *, scales_granularity_mnk: IndexList[Int(3)], transpose_b: Bool = False](c: TileTensor[c_type, Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type], a: TileTensor[a_type, Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type], b: TileTensor[b_type, Engine=b.Engine, address_space=b.address_space, linear_idx_type=b.linear_idx_type], a_scales: TileTensor[a_scales_type, Engine=a_scales.Engine, address_space=a_scales.address_space, linear_idx_type=a_scales.linear_idx_type], b_scales: TileTensor[b_scales_type, Engine=b_scales.Engine, address_space=b_scales.address_space, linear_idx_type=b_scales.linear_idx_type], ctx: DeviceContext)
Computes a batched blockwise scaled FP8 matrix multiplication using a naive per-batch loop that calls the 2D blockwise scaled FP8 kernel for each batch slice.
Parameters:
- c_type (
DType): Output tensor element dtype. - a_type (
DType): LHS input tensor element dtype. - b_type (
DType): RHS input tensor element dtype. - a_scales_type (
DType): LHS scales tensor element dtype. - b_scales_type (
DType): RHS scales tensor element dtype. - scales_granularity_mnk (
IndexList[Int(3)]): Scale granularity(m, n, k); only(1, 128, 128)is currently supported. - transpose_b (
Bool): Whether the RHS input is transposed.
Args:
- c (
TileTensor[c_type, Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type]): Rank-3 output tensor of shape(batch, m, n). - a (
TileTensor[a_type, Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type]): Rank-3 LHS input tensor of shape(batch, m, k). - b (
TileTensor[b_type, Engine=b.Engine, address_space=b.address_space, linear_idx_type=b.linear_idx_type]): Rank-3 RHS input tensor of shape(batch, k, n). - a_scales (
TileTensor[a_scales_type, Engine=a_scales.Engine, address_space=a_scales.address_space, linear_idx_type=a_scales.linear_idx_type]): Rank-3 LHS scales tensor. - b_scales (
TileTensor[b_scales_type, Engine=b_scales.Engine, address_space=b_scales.address_space, linear_idx_type=b_scales.linear_idx_type]): Rank-3 RHS scales tensor. - ctx (
DeviceContext): Device context used to dispatch the per-batch kernels.