linear.py — 这个文件实现了 QATLinear -- 带有假量化(Fake Quantization)功...¶
模块路径: verl.utils.qat.linear
文件概述¶
这个文件实现了 QATLinear -- 带有假量化(Fake Quantization)功能的线性层。"假量化"意味着在前向传播时将权重/激活模拟量化到低精度(FP4),但实际计算仍在原始精度(BF16/FP16)下进行。包含用 Triton 编写的高性能 FP4 假量化内核。
关键代码讲解¶
1. Triton FP4 假量化内核¶
@triton.jit
def _fp4_fake_quant_kernel(x_ptr, y_ptr, M, N, global_scale_ptr, ...):
# 加载一个 tile 的数据
tile = tl.load(x_block_ptr, ...)
# 分组计算 blockwise scale
tile_reshaped = tl.reshape(tile, (TILE_M, NUM_FP4_BLOCKS, BLOCK_SIZE))
block_max = tl.max(tl.abs(tile_reshaped), dim=-1, keep_dims=True)
# 将 scale 量化到 FP8 格式
block_max_quant = block_max_scaled.to(tl.float8e4nv).to(tl.float32) * global_scale
# FP4 量化:将值映射到 {0, 0.5, 1, 1.5, 2, 3, 4, 6} 的 FP4 值
q_val = tl.where(abs_scaled <= 0.25, 0.0,
tl.where(abs_scaled < 0.75, 0.5,
tl.where(abs_scaled <= 1.25, 1.0,
...)))
# 反量化回原始精度
x_rescaled = q_val * block_max_quant_broadcast
FP4 (E2M1) 格式只有 8 个非负值,通过嵌套的 tl.where 实现最近值舍入。
2. 直通估计器(STE)¶
class STEFP4QuantTriton(torch.autograd.Function):
@staticmethod
def forward(ctx, x, global_amax, block_size):
return fp4_fake_quant_weight(x, global_amax=global_amax, block_size=block_size)
@staticmethod
def backward(ctx, grad_output):
return grad_output, None, None # 梯度直接传递,不经过量化
直通估计器(Straight-Through Estimator) 是 QAT 的核心技巧:前向传播执行量化(不可导),反向传播直接传递梯度(假装量化操作是恒等映射)。
3. QATLinear 层¶
class QATLinear(nn.Linear):
def __init__(self, in_features, out_features, bias=True,
mode=QATMode.W4A4, group_size=16, activation_observer="static_minmax", ...):
super().__init__(in_features, out_features, bias, ...)
self.mode = mode
self.group_size = group_size
# W4A4 模式需要额外的 buffer 存储激活的 scale
if mode == QATMode.W4A4:
self.register_buffer("input_global_scale", torch.tensor([-1.0]))
self.register_buffer("input_amax", torch.tensor([-1.0]))
def forward(self, x):
if not self.fake_quant_enabled:
return F.linear(x, self.weight, self.bias)
weight_fq = self._fake_quantize_weight(self.weight) # 量化权重
if self.mode == QATMode.W4A4:
x_fq = self._fake_quantize_activation(x) # 量化激活
else:
x_fq = x
return F.linear(x_fq, weight_fq, self.bias)
4. 权重假量化(带融合)¶
def _fake_quantize_weight(self, weight):
with torch.no_grad():
siblings_ref = getattr(self, "_fusion_siblings_ref", None)
if siblings_ref is not None:
# 融合模式:使用所有兄弟层的最大 amax
all_modules = [self] + siblings
amaxes = [m.weight.abs().max() for m in all_modules]
global_amax = torch.max(torch.stack(amaxes))
else:
global_amax = weight.abs().max()
return STEFP4QuantTriton.apply(weight, global_amax, self.group_size)
5. 激活观测策略¶
def _update_input_global_scale(self, x):
current_amax = torch.amax(torch.abs(x)).detach()
# 多卡同步
if torch.distributed.is_initialized():
torch.distributed.all_reduce(current_amax, op=torch.distributed.ReduceOp.MAX)
if self.activation_observer == "static_minmax":
# 取历史最大值
new_amax = torch.maximum(self.input_amax, current_amax)
elif self.activation_observer == "minmax":
# EMA 移动平均
new_amax = (1 - self._ema_decay) * self.input_amax + self._ema_decay * current_amax
核心类/函数列表¶
| 类/函数名 | 作用 |
|---|---|
QATLinear |
假量化线性层(继承 nn.Linear) |
QATMode |
量化模式枚举 (W4A4/W4A16) |
STEFP4QuantTriton |
直通估计器包装 |
fp4_fake_quant_weight |
FP4 假量化的 Python 入口 |
_fp4_fake_quant_kernel |
Triton 内核实现 |
与其他模块的关系¶
- 被
core.py的apply_qat用来替换nn.Linear - 使用 Triton 编写的高性能 GPU 内核
- 支持 FSDP 分布式训练
小结¶
这是 QAT 模块的计算核心。FP4 假量化通过 Triton 内核实现高性能,STE 解决了量化操作不可导的问题。融合机制(通过 _fusion_siblings_ref)确保 QKV 等层共享 scale。激活观测器在训练过程中动态调整激活值的量化范围。