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 确保每次权重更新后重新计算量化参数。