跳转至

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。激活观测器在训练过程中动态调整激活值的量化范围。