目录

在 INT8 硬件上模拟 MXFP4 矩阵乘

(本文由 GLM-5.2 生成,仅供参考,反正我也不发知乎)

DeepSeek V4 等新一代 MoE 大模型,路由专家(Router Experts)的权重广泛采用 MXFP4 格式以压缩体积、降低访存。而很多硬件最低只支持 INT8 矩阵乘,没有 FP8/FP4 算力。若能在 INT8 硬件上模拟 MXFP4 矩阵乘,就能同时拿到 4-bit 权重的访存优势和 INT8 的算力优势。

MXFP4 数据以 FP4 存储,采用 Per-Block 量化:32 个数共享一个缩放因子,类型 FP8_ue8m0(8 位无符号指数),值恰为 2 的整数幂。

# 1. Per-Channel 量化

INT8 GEMM 依赖 Per-Channel 量化(一个通道共享一个缩放因子):通道内缩放相同,就能先做纯整数累加(INT8 × INT8 → INT32),最后再转浮点乘缩放。Per-Block 量化每个 block 各自有缩放,累加中途就得转浮点,性能很差。所以核心问题是把 Per-Block 量化改造成 Per-Channel 量化

# 2. 方案概述

FP4 的 16 个值乘以 16 恰好都能用 INT8 无损表示:

[0.0,0.5,1,1.5,2,3,4,6,0.0,0.5,1,1.5,2,3,4,6]×16[0.0,\ 0.5,\ 1,\ 1.5,\ 2,\ 3,\ 4,\ 6,\ -0.0,\ -0.5,\ -1,\ -1.5,\ -2,\ -3,\ -4,\ -6] \times 16

=[0,8,16,24,32,48,64,96,0,8,16,24,32,48,64,96]=[0,\ 8,\ 16,\ 24,\ 32,\ 48,\ 64,\ 96,\ 0,\ -8,\ -16,\ -24,\ -32,\ -48,\ -64,\ -96]

于是用查表指令把 4-bit FP4 映射成 INT8。又因 MXFP4 缩放因子是 2 的幂,统一 block 之间的缩放就是右移:各 block 权重右移对齐到同一尺度后跑 INT8 GEMM 累加,最后乘一个 Per-Channel 缩放因子。

转换是在线的:权重以 MXFP4 存储,每个 block 都要右移再喂给矩阵乘指令,两者交替出现。但是右移会丢精度,需要一定的预处理。

# 3. 舍入处理

目标是向最近 / 偶数舍入(RNE)。

由于是权重,预处理阶段已经知道每个元素的 INT8 值 nn 和右移位数 bb,所以 RNE 的取舍可以完全离线算好,直接改写存储的 FP4 值(预舍入)。在线计算还是右移 v = n >> b

例如 n = 24, b = 4,预处理把 n 预舍入为 32,在线右移的结果即为 2。2424=1.5\frac{24}{2^4}=1.5 向偶数舍入也是 2,这样就能对上。

如果 n = 96, b = 6,预处理把 n 预舍入为 128,但是 128 不在 FP4 × 16 的表里,只能让整个 Block 的 b 减少 1 做兼容。

const int8_t fp4_to_int8[16] = {
    0,  8, 16, 24, 32, 48, 64, 96,
    0, -8,-16,-24,-32,-48,-64,-96};

int8_t shift_weight(unsigned fp4, int b) {
    return fp4_to_int8[fp4] >> b;
}

# 4. 完整实施方案

预处理(离线做一次)做两件事:

  1. 拆缩放因子:把每个通道的 MXFP4 缩放因子拆成 Per-Block 右移位数 b=emaxeb0b = e_{\max} - e_b \ge 0(取通道内最大缩放因子 2emax2^{e_{\max}} 为基准,都是 2 的幂,除法等价于指数相减)与 Per-Channel 缩放因子 2emax2^{e_{\max}}(通道内唯一,留给最后浮点缩放)。
  2. 预舍入:改写每个元素的 n 与块缩放因子 b,使在线右移结果等于 RNE。

在线计算的指令比较简单:

  1. 加载 fp4 值和 b,查表得到 n,计算 n >> b
  2. 利用矩阵乘指令,进行 INT8 乘法并累加结果。
  3. 完整通道算完后,乘以 Per-Channel 缩放因子。