看到Kimi-K3中使用了量化感知训练(QAT),这里记录一下在 PyTorch
中如何实现 QAT。
在 PyTorch 中实现 QAT
PyTorch 提供了一个专门的 torch.quantization 模块来方便
QAT。一般流程包含为 QAT 准备模型、进行微调
(fine-tuning),然后将其转换为真正的量化 (quantization)模型。
模型准备
- 定义一个
QConfig,它指定量化设置(例如,用于激活统计的观测器、用于权重
(weight)和激活的伪量化模块、目标数据类型如
torch.qint8)。
- 在想要量化的模型部分的开头和结尾插入
QuantStub 和
DeQuantStub 层。这些层作为标记
(token),告知框架量化操作的起点和终点。
- 使用
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
float_model = MyModel() float_model.train()
qconfig = torch.quantization.get_default_qat_qconfig('fbgemm')
prepared_model = torch.quantization.prepare_qat(float_model, {'': qconfig})
print(prepared_model)
|
微调
- 像训练常规浮点模型一样训练
prepared_model。使用您的标准训练循环、损失函数 (loss
function)和优化器。
- 在前向传播过程中,伪量化模块根据观测器收集的统计数据模拟量化效果(钳制、舍入)。
- 在反向传播 (backpropagation)过程中,STE
允许梯度通过模拟量化步骤回传,使模型权重能够适应量化过程。
- 通常做法是从预训练 (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() optimizer.step()
|
转换为量化模型
- 微调后,将模型切换到评估模式(
prepared_model.eval())。
- 使用
torch.quantization.convert 将经过 QAT
训练的模型转换为真正的量化模型。这会使用学到的参数
(parameter),将伪量化模块和观测到的浮点模块(如
nn.Linear)替换为它们的基于整数的对应模块(如
nn.quantized.Linear)。
1 2 3 4 5 6 7 8 9 10
| prepared_model.eval()
quantized_model = torch.quantization.convert(prepared_model.cpu())
print(quantized_model)
|