面向输出头量化的 Softmax 重参数化

  • 关联论文:2609.31291
  • 作者:Tom
  • 更新:2026-09-29

一句话结论

大词表使输出头成为小型语言模型推理成本的重要组成部分,而量化输出头会显著扭曲 softmax 分布。Softmax 重参数化通过在量化前在输出头的加性等价类中搜索最优代表——对每个输出行减去词表行均值的标量倍数,以验证集 KL 为准则选取系数——在 W4 量化下实现 test KL 下降 73~93%,Phi 模型首批生成延迟降低 10.8%,且对线性 softmax head 保持全精度预测完全不变、推理零额外开销。

解决什么真问题

小型语言模型(SLM)通常配置较大的词表(如 100K+ 词元),而输出投影层(output projection / output head)包含 V×d 个权重(V 为词表大小,d 为隐层维度),每次解码步骤都必须读取完整的 2V d 字节权重数据。由于单 token 解码是内存带宽受限(memory-bound)的操作,这一大量权重移动直接转化为推理延迟。

当前实践中,输出头往往保留在高精度(BF16),而解码器其他部分接受低比特量化。这导致一个悖论:即使解码器被压缩到 W4,输出头仍以 BF16 运行,全精度读取 2V d 字节的权重成为内存瓶颈。

核心张力: - 输出头直接决定 next-token 概率,没有后续学习层吸收量化误差 - Softmax 只依赖相对 logit,不依赖绝对值:相同 softmax 分布可以由不同的权重矩阵实现 - 不同等价权重矩阵的量化残差不同:这带来了选择自由

换言之,同一个全精度输出头可以用无数个「功能等价」的权重矩阵表示(加性等价类),它们产生完全相同的 softmax 分布,但对量化的友好程度不同。Softmax 重参数化正是利用这一自由度,在量化前选取对量化最友好的那个代表。

核心方法

机制拆解

加性等价类(Additive Equivalence Class):

设输出头权重矩阵为 $W \in \mathbb{R}^{V \times d}$,输出 logit 为 $o = xW$(x 为隐状态)。对于任意标量 $\alpha$,定义:

$$W^{(\alpha)} = W - \alpha \cdot \mathbf{1}\boldsymbol{\mu}^\top$$

其中 $\boldsymbol{\mu} \in \mathbb{R}^d$ 是词表行的均值向量,$\mathbf{1} \in \mathbb{R}^V$ 是全 1 向量。$W^{(\alpha)}$ 和 $W$ 对任意输入 x 产生相同的 softmax 分布,因为:

$$o^{(\alpha)} = xW^{(\alpha)} = xW - \alpha \cdot x\mathbf{1}\boldsymbol{\mu}^\top = o - \alpha (\mathbf{1}^\top x) \boldsymbol{\mu}^\top$$

对所有 logit 同时减去同一个常数($\alpha$ 乘以某个标量),Softmax 对常数平移不变,因此 $p = \text{softmax}(o) = \text{softmax}(o^{(\alpha)})$。

核心方法:

对每个 $\alpha$ 值,量化 $W^{(\alpha)}$ 至目标精度(如 W4),在验证集上测量量化后的 softmax 分布与全精度分布的 KL 散度,选取 KL 最小的 $\alpha^*$:

$$\alpha^* = \arg\min_{\alpha} \text{KL}\left( p_{\text{fp}} \parallel p_{\text{quant}}(W^{(\alpha)}) \right)$$

这是一个一维搜索,计算代价极低。

线性 softmax head 的精确保持:

对于线性 softmax head(标准 LLM 输出),$W^{(\alpha^)}$ 与原始 $W$ 产生完全相同的全精度 softmax 分布(Softmax 对加性常数的平移不变性是精确的),因此不需要重新训练解码器*,这是一次性预处理。

非线性 logit 路径的扩展:

部分 LLM 在 LM head 前有非线性变换(如某些 attention 变体)。对于这类路径,文章提出了一个 rank-1 修正来扩展精确保持范围。

量化器兼容性:

方法兼容多种后训练量化方法:RTN(Round-To-Nearest)、activation-weighted MSE、Full-Hessian GPTQ。对每种量化器单独进行 $\alpha$ 搜索,因为不同的量化器对权重分布的敏感度不同。

关键公式

# Softmax 重参数化算法(伪代码)
for each output head:
    μ = mean(W, axis=0)           # 词表行均值向量
    best_alpha = 0
    best_kl = inf

    for alpha in search_grid:      # 一维网格搜索
        W_shifted = W - alpha * 𝟙 * μ^T
        W_quant = quantize(W_shifted, bits=W)
        p_quant = softmax(x @ W_quant)  # 在验证集上
        p_full  = softmax(x @ W)       # 全精度参考
        kl = KL(p_full || p_quant)
        if kl < best_kl:
            best_kl = kl
            best_alpha = alpha

    W_final = quantize(W - best_alpha * 𝟙 * μ^T, bits=W)
    # 存储 W_final 和 best_alpha(解码时需要还原... 实际上对 shift-compatible heads shift 不改变 logit 相对顺序,可以吸收到 softmax 实现中)

⚠️ 存疑:shift 的实现细节(如何在量化后还原或吸收 shift)原文有专门讨论,但此处未完全明确,以下描述基于 abstract 推断。

关键实验与数据

  • 模型:XGLM、Phi、Phi-2、BLOOM、BLOOMZ 等 7 个输出头
  • 量化器:RTN、activation-weighted MSE、Full-Hessian GPTQ
  • 验证数据集:WikiText、C4、OpenWebMath(跨域泛化验证)
模型 量化精度 量化器 test KL 下降幅度
XGLM W4 RTN 93%
Phi W4 activation-weighted MSE 73%
BLOOM W4 activation-weighted MSE 73~77%
BLOOMZ W4 activation-weighted MSE 73~77%
  • 更强 GPTQ 校准下:Phi 上的收益得以保持
  • 跨域泛化:在 WikiText 上选择的 $\alpha^*$ 无需重新调整即可迁移到 C4 和 OpenWebMath
  • 延迟收益:decoder 保持 BF16、output head 压缩至 W4 的 Phi,首批生成延迟降低 10.8%(batch-one)
  • 对低误差 head:baseline 量化误差已较低的 head,重参数化后改善不明显

亮点与局限

亮点: - 理论上优雅:利用 softmax 对加性平移的不变性,将重参数化精确控制在 softmax 等价类内,保证全精度预测完全不变 - 零推理开销(对 shift-compatible heads):shift 不改变 logit 相对顺序,可以融入 softmax 实现中,不引入额外矩阵乘法 - 跨模型、跨量化器通用:7 个不同模型 + 3 种量化器均验证有效 - 跨域泛化能力强:WikiText 选择的系数直接迁移到 C4 和 OpenWebMath,无需 retune - 延迟收益有实际意义:10.8% 的首批生成延迟改善在工程上有价值

局限: - 仅在 small language models 上验证,对 GPT-4 级别的大模型泛化性未测(但大模型输出头量化可能更关键,因为内存压力更大) - 对 tied embedding(输入 embedding 与输出 head 共享权重)的处理有额外约束,论文有讨论但此处未深入 - rank-1 修正扩展到非线性路径的效果和通用性尚需更多验证 - 论文为 ICLR 2027 投稿,尚未经顶会评审

对工程落地的启发

  1. 输出头量化是小模型推理优化的重要被忽视角落:当解码器已压缩至 W4 时,输出头仍是 BF16 的内存带宽瓶颈,重参数化提供了低成本的补齐方案。
  2. 加性等价类是量化友好的隐式优化空间:对其他受限于输出分布的组件(如某些 attention 变体),类似的重参数化思路可能有推广价值。
  3. 验证集 KL 导向的参数搜索对量化友好:不需要重建训练流程,只需一次离线搜索即可,适用于无法重新训练的场景。
  4. 系数跨数据集泛化的意义:在领域相关数据上做一次搜索,可以在多个相关领域通用,降低了工程部署成本。
  5. Softmax 平移不变性作为理论基石:对想深入理解量化和知识蒸馏关系的研究者,这是一个值得关注的理论基础。

与同方向工作的关系

工作 方法 与本文关系
Channel Scaling(Xiao et al., 2023) 通道缩放 同属量化友好型重参数化,但针对激活,本文针对输出头
ASR + Mean-Centering(SageAttention) 对 K 做均值中心化 固定均值中心化;本文是针对输出的可学习一维搜索
Output Embedding Centering(Stollenwerk et al., 2026) 预训练阶段词表行均值中心化 预训练介入;本文是后训练一次性处理
GPTQ Hessian-based 量化 本文可与 GPTQ 结合,为 GPTQ 选择最优输出头代表

本文填补了「输出头量化」这一小模型推理优化的空白,是 LLM 推理优化方向的重要补充。

适合谁读

  • LLM 推理优化 / 模型压缩方向的研究者和工程师:尤其是做 INT4 / W4 量化、关注解码延迟优化的团队
  • 后训练量化(PTQ)方向研究者:对量化友好的重参数化技术有兴趣的读者
  • 小模型(SLM)实践者:Phi / SmolLM / Granite 等模型的部署优化
  • 对 softmax 理论性质有研究兴趣的学者:关注加性等价类与量化误差关系的理论工作

⚠️ 存疑:非线性 logit 路径的 rank-1 修正方案的具体实现和适用范围原文有专门讨论,本解读基于 abstract 推断;ICLR 2027 投稿,尚未经顶会评审;系数跨域泛化的理论保证尚待进一步验证;shift 在量化推理引擎(如 vLLM / AWQ)中的实际集成方式需参考原文字代码。