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
convert_ref_scales_to_mxfp8_format
def convert_ref_scales_to_mxfp8_format[MType: CoordLike, NType: CoordLike, KType: CoordLike, ref_a_scales_tt_layout: TensorLayout, ref_b_scales_tt_layout: TensorLayout, a_scales_tt_layout: TensorLayout, b_scales_tt_layout: TensorLayout, //, ref_scales_type: DType, scales_type: DType, *, REF_BLOCK_SIZE: Int, SF_VECTOR_SIZE: Int](m: MType, n: NType, k: KType, ref_a_scales: TileTensor[ref_scales_type, ref_a_scales_tt_layout, Engine=ref_a_scales.Engine, linear_idx_type=ref_a_scales.linear_idx_type], ref_b_scales: TileTensor[ref_scales_type, ref_b_scales_tt_layout, Engine=ref_b_scales.Engine, linear_idx_type=ref_b_scales.linear_idx_type], a_scales: TileTensor[scales_type, a_scales_tt_layout, Engine=a_scales.Engine, linear_idx_type=a_scales.linear_idx_type], b_scales: TileTensor[scales_type, b_scales_tt_layout, Engine=b_scales.Engine, linear_idx_type=b_scales.linear_idx_type])
Converts reference float32 block scales into the 5D MXFP8 E8M0 scale-factor layout.
Reads the per-block float32 reference scales for the A (M x K) and
B (N x K) operands, converts each to float8_e8m0fnu, and writes them into
the corresponding 5D scale-factor tensors expected by block-scaled matmul
kernels.
Parameters:
- MType (
CoordLike): CoordLike type carrying the M dimension size. - NType (
CoordLike): CoordLike type carrying the N dimension size. - KType (
CoordLike): CoordLike type carrying the K dimension size. - ref_a_scales_tt_layout (
TensorLayout):TensorLayoutof the 2D reference A scales. - ref_b_scales_tt_layout (
TensorLayout):TensorLayoutof the 2D reference B scales. - a_scales_tt_layout (
TensorLayout):TensorLayoutof the 5D output A scales. - b_scales_tt_layout (
TensorLayout):TensorLayoutof the 5D output B scales. - ref_scales_type (
DType): Element type of the reference scales (must be float32). - scales_type (
DType): Element type of the output scales (must be float8_e8m0fnu). - REF_BLOCK_SIZE (
Int): Block size (in elements) used by the reference scales. - SF_VECTOR_SIZE (
Int): Number of elements each scale factor covers in the output layout.
Args:
- m (
MType): M dimension size of the operands. - n (
NType): N dimension size of the operands. - k (
KType): K dimension size of the operands. - ref_a_scales (
TileTensor[ref_scales_type, ref_a_scales_tt_layout, Engine=ref_a_scales.Engine, linear_idx_type=ref_a_scales.linear_idx_type]): 2D float32 reference scales for the A operand, indexed as[k // REF_BLOCK_SIZE, m]. - ref_b_scales (
TileTensor[ref_scales_type, ref_b_scales_tt_layout, Engine=ref_b_scales.Engine, linear_idx_type=ref_b_scales.linear_idx_type]): 2D float32 reference scales for the B operand, indexed as[n // REF_BLOCK_SIZE, k // REF_BLOCK_SIZE]. - a_scales (
TileTensor[scales_type, a_scales_tt_layout, Engine=a_scales.Engine, linear_idx_type=a_scales.linear_idx_type]): Mutable 5D output tensor receiving the converted A scales. - b_scales (
TileTensor[scales_type, b_scales_tt_layout, Engine=b_scales.Engine, linear_idx_type=b_scales.linear_idx_type]): Mutable 5D output tensor receiving the converted B scales.