、全连接(Linear)、池化(Pooling)和 LSTM,而没有 ReLU/SiLU 等激活函数)
Gemmi 生成# 将一个普通的 PyTorch 神经网络模块如 Conv2d“变身”为一个量化感知模块如 QuantConv2d而不丢失原有模块的所有参数和状态。deftransfer_torch_to_quantization(nninstance:torch.nn.Module,quantmodule):# 创建未初始化的量化实例quant_instancequantmodule.__new__(quantmodule)# “偷梁换柱”复制原有属性fork,valinvars(nninstance).items():setattr(quant_instance,k,val)# 此时quant_instance 已经有了正确的权重和偏置但它还没有加载量化器。它本质上是一个“披着量化类外衣的空壳”。def__init__(self):# 从全局配置或类定义中提取输入激活值和权重权重参数的量化描述符如校准方法 max位宽 8 等。quant_desc_input,quant_desc_weightquant_nn_utils.pop_quant_desc_in_kwargs(self.__class__)# 仅量化输入 (如某些激活函数或特殊层)ifisinstance(self,quant_nn_utils.QuantInputMixin):#quant_desc_input quant_nn_utils.pop_quant_desc_in_kwargs(self.__class__, input_onlyTrue)self.init_quantizer(quant_desc_input)# Turn on torch_hist to enable higher calibration speedsifisinstance(self._input_quantizer._calibrator,calib.MaxCalibrator):self._input_quantizer._calibrator._torch_histTrueelse:# 量化输入和权重 (如标准的卷积层)self.init_quantizer(quant_desc_input,quant_desc_weight)# Turn on torch_hist to enable higher calibration speedsifisinstance(self._input_quantizer._calibrator,calib.MaxCalibrator):self._input_quantizer._calibrator._torch_histTrueself._weight_quantizer._calibrator._torch_histTrue__init__(quant_instance)returnquant_instancedefreplace_to_quantization_module(model:torch.nn.Module):module_dict{}# 遍历输入的 PyTorch 模型找到那些被注册在 quant_modules._DEFAULT_QUANT_MAP 中的特定层print(quant_modules:,quant_modules._DEFAULT_QUANT_MAP)forentryinquant_modules._DEFAULT_QUANT_MAP:# 原始模块所在的类或模块例如 mmcv.cnn.bricks.wrappers--原始模块的具体类名字符串例如 ConvTranspose2dmodulegetattr(entry.orig_mod,entry.mod_name)# 用于替换的量化模块类例如 quant_nn.QuantConvTranspose2d。# id(module): 获取该类对象在内存中的唯一标识符。module_dict[id(module)]entry.replace_mod# 将它们替换为对应的量化感知模块# 这是一个深度优先搜索DFS函数用于遍历整个模型树。defrecursive_and_replace_module(module,prefix):fornameinmodule._modules:submodulemodule._modules[name]pathnameifprefixelseprefix.name recursive_and_replace_module(submodule,path)submodule_idid(type(submodule))ifsubmodule_idinmodule_dict:print(submodule_id:,submodule_id,module_dict[submodule_id])module._modules[name]transfer_torch_to_quantization(submodule,module_dict[submodule_id])recursive_and_replace_module(model)quant_modules: [ quant_entry( orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConv1d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConv1d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConv2d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConv2d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConv3d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConv3d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConvTranspose1d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConvTranspose1d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConvTranspose2d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConvTranspose2d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameConvTranspose3d, replace_modclass pytorch_quantization.nn.modules.quant_conv.QuantConvTranspose3d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameLinear, replace_modclass pytorch_quantization.nn.modules.quant_linear.QuantLinear), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameLSTM, replace_modclass pytorch_quantization.nn.modules.quant_rnn.QuantLSTM), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameLSTMCell, replace_modclass pytorch_quantization.nn.modules.quant_rnn.QuantLSTMCell), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAvgPool1d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAvgPool1d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAvgPool2d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAvgPool2d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAvgPool3d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAvgPool3d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAdaptiveAvgPool1d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAdaptiveAvgPool1d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAdaptiveAvgPool2d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAdaptiveAvgPool2d), quant_entry(orig_modmodule torch.nn from /torch/nn/__init__.py, mod_nameAdaptiveAvgPool3d, replace_modclass pytorch_quantization.nn.modules.quant_pooling.QuantAdaptiveAvgPool3d)] 为何只有卷积没有激活函数这个问题问得很到点子上你在日志里看到的这个quant_modules映射表实际上是pytorch_quantization的猴子补丁机制Monkey Patching /quant_modules.initialize()在工作。之所以这个列表中主要只有卷积Conv、全连接Linear、池化Pooling和 LSTM而没有 ReLU/SiLU 等激活函数主要有以下两个核心原因1. 硬件部署视角激活函数的量化节点会被“融合Fusion”在 TensorRT 或 CUDA 硬件加速推理时ReLU / GELU / SiLU 等激活函数通常不会作为独立的硬件指令执行。层融合Layer FusionTensorRT 在编译 Engine 时会自动把Conv Bias ReLU融合成一个单一的 CUDA 内核Kernel。截断自动完成卷积层输出时带有的量化参数Scale / Zero-point已经能够把负数截断ReLU 的作用就是把000的部分清零激活函数的逻辑已经在硬件计算卷积的算子末端顺带完成了。无需额外伪量化FakeQuant如果给 ReLU 节点单独插入一个伪量化算子不仅无法提升精度反而会在 TensorRT 导出时引入多余的 Tensor 转换开销甚至破坏 TensorRT 的算子融合优化。2. 量化算法视角计算密集型 vs 激活/非线性节点pytorch_quantization的设计初衷是做INT8 权重量化Weight Quantization与激活值量化Activation Quantization重点监控的是权重/输入计算量极大的算子算子类型代表模块为什么需要替换为Quantized版本计算密集型Conv2d,Linear包含可学习参数权重Weight需要为其插入Weight FakeQuantizer并在输入端插入Input FakeQuantizer。下采样/汇聚型AvgPool2d,AdaptiveAvgPool2d影响特征图数值范围分布且涉及张量重组需要独立的Input/Output FakeQuantizer来对齐 Scale。非线性激活ReLU,SiLU,Sigmoid无权重仅对输入元素做逐元素Element-wise映射。在标准 INT8 量化体系中它的输入/输出 Scale 可以直接继承前一级 Conv 或下一级 Conv 的 Scale无需单独替换 Class。3. 那激活值的量化Activation Quantization是怎么做的虽然没有QuantReLU这个类但这并不意味着激活值没有被量化当你调用quant_modules.initialize()时Conv2d被替换成了QuantConv2d。在QuantConv2d内部包含了两个量化器weight_quantizer专门用来伪量化卷积核权重。input_quantizer专门用来伪量化传入该卷积层的输入特征图Activation。也就是说激活值的量化已经被打包收纳在下一个层比如下一个 Conv 或 Linear的输入端或者上一个层的输出端了根本不需要额外去替换torch.nn.ReLU。 总结与建议你看到的映射表只是pytorch_quantization自动将 PyTorch 默认带有权重的算子替换为支持伪量化Fake Quantization插桩的子类。对于 MMDetection/MMRazor 训练这种设计完全符合 NVIDIA 官方推荐的 INT8 量化范式。在校准Calibration或 QAT 微调时直接按照正常流程启用quant_modules.initialize()即可无需手动给激活函数添加量化节点。