跳转至

core.py — 这是 QAT 模块的核心配置和应用逻辑

模块路径: verl.utils.qat.core

文件概述

这是 QAT 模块的核心配置和应用逻辑。主要功能包括:QAT 配置管理(QATConfig)、将模型中的 nn.Linear 替换为 QATLinear(apply_qat)、设置权重 scale 融合(enable_qat_fuse)。

关键代码讲解

1. QAT 配置类

@dataclass
class QATConfig(BaseConfig):
    enable: bool = False                # 是否启用 QAT
    mode: str = "w4a16"                 # 量化模式: "w4a4" 或 "w4a16"
    group_size: int = 16                # 量化分组大小
    ignore_patterns: list[str] = field(  # 不做量化的层
        default_factory=lambda: ["lm_head", "embed_tokens", "re:.*mlp.gate$"]
    )
    activation_observer: str = "static_minmax"  # 激活值观测策略
    quantization_config_path: Optional[str] = None  # 量化配置文件路径

ignore_patterns 支持普通字符串匹配和正则表达式(以 re: 开头)。通常 embedding 层和输出层不做量化。

2. 判断是否需要量化

def _should_quantize(name: str, module: nn.Module, config: QATConfig) -> bool:
    if not isinstance(module, nn.Linear):
        return False                    # 只量化 Linear 层
    for pattern in config.ignore_patterns:
        if pattern.startswith("re:"):
            if re.match(pattern[3:], name):  # 正则匹配
                return False
        else:
            if pattern in name:              # 子串匹配
                return False
    if module.in_features % config.group_size != 0:
        return False                    # 输入维度必须能被 group_size 整除
    return True

3. 应用 QAT

def apply_qat(model, config):
    """将模型中的 nn.Linear 替换为 QATLinear"""
    mode = QATMode(config.mode.lower())
    modules_to_replace = []
    for name, module in model.named_modules():
        if _should_quantize(name, module, config):
            modules_to_replace.append((name, module))

    for name, module in modules_to_replace:
        fake_quant_module = QATLinear.from_linear(
            module, mode=mode, group_size=config.group_size,
            activation_observer=config.activation_observer
        )
        _set_module(model, name, fake_quant_module)
    return model

4. 权重 Scale 融合

FUSION_PATTERNS = {
    "qkv": ["q_proj", "k_proj", "v_proj"],
    "gate_up": ["gate_proj", "up_proj"],
}

def setup_fusion_siblings(model):
    """为 QKV 和 GateUp 层设置融合兄弟关系"""
    # QKV 的三个投影使用相同的权重 scale,以减少量化误差
    for parent, projs in groups.items():
        modules = list(projs.values())
        for i, m in enumerate(modules):
            siblings = modules[:i] + modules[i+1:]
            m._fusion_siblings_ref = [weakref.ref(s) for s in siblings]

融合的目的:Transformer 中的 Q/K/V 投影通常会被融合为一个矩阵乘法,它们共享相同的量化 scale 可以减少误差。

5. 清除缓存的 Scale

def invalidate_all_scales(model):
    """在 optimizer.step() 后调用,清除缓存的权重 scale"""
    for module in model.modules():
        if isinstance(module, QATLinear):
            module._weight_blockwise_scale = None
            module._weight_global_scale = None
            module._cached_weight_amax = None

核心类/函数列表

类/函数名 作用
QATConfig QAT 配置数据类
apply_qat 将 nn.Linear 替换为 QATLinear
enable_qat_fuse 启用融合模式
invalidate_all_scales 清除缓存的量化 scale
load_quantization_config 加载量化配置 JSON
setup_fusion_siblings 设置 QKV/GateUp 的融合兄弟

与其他模块的关系

  • 导入 linear.py 的 QATLinear 和 QATMode
  • 被 __init__.py 导出
  • 在 FSDP 包裹模型之前调用 apply_qat

小结

这是 QAT 的配置和应用入口。核心操作是遍历模型的所有 nn.Linear,将符合条件的替换为 QATLinear(假量化版本)。融合机制确保 QKV 等相关层共享量化参数,invalidate_all_scales 确保每次权重更新后重新计算量化参数。