This document is relevant for: Trn2, Trn3
nki.isa.quantize_mx#
- nki.isa.quantize_mx(dst, src, dst_scale, name=None)[source]#
Quantize FP16/BF16 data to MXFP8 tensors (both data and scales) using Vector Engine.
Note
Available only on NeuronCore-v4 and newer.
The resulting
dstanddst_scaletensors use the MXFP8 element and scale data types as defined in the OCP Microscaling standard. This instruction calculates the required scales for each group of 32 values insrc, divides them by the calculated scale, and casts to the target MXFP8 element data type.The scale calculation differs from the sample conversion algorithm described in the OCP specification: this instruction uses a block scale that is two times larger, reserving additional range for rounding the largest values in each group without saturation. This remains OCP MX-compliant.
The output layout is suitable for direct consumption by the
nisa.nc_matmul_mxAPI running on Tensor Engine.Memory types.
All input
srcand output tiles (dstanddst_scale) must be in SBUF.Data types.
The input
srctile must be float16 or bfloat16. The outputdsttile must be float8_e5m2_x4 or float8_e4m3fn_x4 (4-packed FP8 data types). Thedst_scaletile must be float8_e8m0fnu or uint8 (preferfloat8_e8m0fnu: OCP MX standard; uint8 accepted for backward compatibility).The 4-packed data types (float8_e5m2_x4/float8_e4m3fn_x4) are 32-bit data types that pack four 8-bit float8_e5m2/float8_e4m3fn values.
Layout.
The quantization operates on groups of 32 elements from the input
srctile, where each group consists of 8 partitions × 4 elements per partition. For each 32-element group, the instruction produces:Quantized FP8 data in
dstOne shared scale value in
dst_scaleper group
Tile size.
The partition dimension size of
srcmust be a multiple of 32 and must not exceed 128.The free dimension size of
srcmust be a multiple of 4 and must not exceed the physical size of each SBUF partition.The
dsttile has the same partition dimension size assrcbut a free dimension size that is 1/4 ofsrcfree dimension size due to the special 4-packed FP8 data types.
Scale calculation.
For a group \(V\) of 32 values, let
\[a_{\max} = \max_{V_i \in V} |V_i|\]Let \(E_{\max}\) be the maximum unbiased exponent of the destination element data type: 8 for
float8_e4m3fnand 15 forfloat8_e5m2.The block scale \(X\) is calculated as
\[X = 2^{ \left\lfloor \log_2(a_{\max}) \right\rfloor - (E_{\max} - 1) }\]For an all-zero group, where \(a_{\max} = 0\), the block scale is set to \(X = 2^{-127}\), the minimum value representable by
float8_e8m0fnu.- Parameters:
dst – the quantized MXFP8 output tile
src – the input FP16/BF16 tile to be quantized
dst_scale – the MXFP8 output scale tile (float8_e8m0fnu or uint8)
This document is relevant for: Trn2, Trn3