量化感知训练(Quantization-aware Training, QAT)

看到Kimi-K3中使用了量化感知训练(QAT),这里记录一下在 PyTorch 中如何实现 QAT。

在 PyTorch 中实现 QAT

PyTorch 提供了一个专门的 torch.quantization 模块来方便 QAT。一般流程包含为 QAT 准备模型、进行微调 (fine-tuning),然后将其转换为真正的量化 (quantization)模型。

模型准备

  1. 定义一个QConfig,它指定量化设置(例如,用于激活统计的观测器、用于权重 (weight)和激活的伪量化模块、目标数据类型如 torch.qint8)。
  2. 在想要量化的模型部分的开头和结尾插入 QuantStubDeQuantStub 层。这些层作为标记 (token),告知框架量化操作的起点和终点。
  3. 使用 torch.quantization.prepare_qat 根据提供的 QConfig 自动将伪量化模块和观测器插入到您的模型中。此函数会就地修改模型,或返回一个为 QAT 准备好的新模型实例。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
import torch
import torch.nn as nn
import torch.quantization

# 示例:一个简单模型
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.quant = torch.quantization.QuantStub() # 输入量化标记
self.linear = nn.Linear(10, 20)
self.relu = nn.ReLU()
self.dequant = torch.quantization.DeQuantStub() # 输出去量化标记

def forward(self, x):
x = self.quant(x) # 对输入应用伪量化
x = self.linear(x)
x = self.relu(x)
x = self.dequant(x) # 返回浮点数前对输出去量化
return x

# 1. 实例化浮点模型
float_model = MyModel()
float_model.train() # 将模型设置为训练模式以进行 QAT

# 2. 定义 QConfig(支持 INT8 对称逐张量后端示例)
# 根据目标硬件/后端需求进行调整。
qconfig = torch.quantization.get_default_qat_qconfig('fbgemm') # 或 'qnnpack' 等

# 3. 为 QAT 准备模型
prepared_model = torch.quantization.prepare_qat(float_model, {'': qconfig})

print(prepared_model) # 观察插入的 FakeQuantize 模块

微调

  1. 像训练常规浮点模型一样训练 prepared_model。使用您的标准训练循环、损失函数 (loss function)和优化器。
  2. 在前向传播过程中,伪量化模块根据观测器收集的统计数据模拟量化效果(钳制、舍入)。
  3. 在反向传播 (backpropagation)过程中,STE 允许梯度通过模拟量化步骤回传,使模型权重能够适应量化过程。
  4. 通常做法是从预训练 (pre-training)的浮点模型检查点开始 QAT,并以较小的学习率微调几个周期。
1
2
3
4
5
6
7
8
9
10
11
num_epochs_qat = 3 # 通常比初始训练的周期数少

for epoch in range(num_epochs_qat):
prepared_model.train() # 确保模型处于训练模式
for data, target in train_loader:
optimizer.zero_grad()
output = prepared_model(data)
loss = criterion(output, target)
loss.backward() # 梯度通过 STE 流经伪量化节点
optimizer.step()
# 如有需要,添加验证循环

转换为量化模型

  1. 微调后,将模型切换到评估模式(prepared_model.eval())。
  2. 使用 torch.quantization.convert 将经过 QAT 训练的模型转换为真正的量化模型。这会使用学到的参数 (parameter),将伪量化模块和观测到的浮点模块(如 nn.Linear)替换为它们的基于整数的对应模块(如 nn.quantized.Linear)。
1
2
3
4
5
6
7
8
9
10
# 转换前确保模型处于评估模式
prepared_model.eval()

# 将 QAT 模型转换为可部署的量化模型
quantized_model = torch.quantization.convert(prepared_model.cpu()) # 通常先转换为 CPU 模型

print(quantized_model) # 观察量化模块(例如,QuantizedLinear)

# 现在 'quantized_model' 可以保存并用于推理
# torch.save(quantized_model.state_dict(), "quantized_model.pth")