版本确认 
horizon-plugin-pytorch 2.3.6 是比较新的版本,基于 PyTorch 1.13。这个版本下,手动绑定 observer 的方式可能因 API 变化而失效。
针对 2.3.6 版本的解决方案
方案 1:使用 torch.ao.quantization 的 share_qconfig 机制(推荐)
2.3.6 版本可能支持通过 qconfig_dict 的特殊语法来共享 observer:
import torch
import torch.nn as nn
import copy
from horizon_plugin_pytorch.quantization import (
prepare_qat_fx,
convert_fx,
set_fake_quantize,
FakeQuantState,
QuantStub,
)
from horizon_plugin_pytorch.quantization.qconfig import default_qat_8bit_fake_quant_qconfig
class MultiScaleCatModule(nn.Module):
def __init__(self, in_channels=128):
super().__init__()
# 为每个输入路径创建独立的 QuantStub
self.quant_k_0 = QuantStub()
self.quant_k_1 = QuantStub()
self.quant_k_2 = QuantStub()
self.quant_v_0 = QuantStub()
self.quant_v_1 = QuantStub()
self.quant_v_2 = QuantStub()
self.conv_0 = nn.Conv2d(in_channels, in_channels, 3, padding=1)
self.conv_1 = nn.Conv2d(in_channels, in_channels, 3, padding=1)
self.conv_2 = nn.Conv2d(in_channels, in_channels, 3, padding=1)
def forward(self, x):
k_0 = self.conv_0(x)
k_1 = self.conv_1(x)
k_2 = self.conv_2(x)
v_0 = k_0
v_1 = k_1
v_2 = k_2
k_features = torch.cat([
self.quant_k_0(k_0),
self.quant_k_1(k_1),
self.quant_k_2(k_2),
], dim=1)
v_features = torch.cat([
self.quant_v_0(v_0),
self.quant_v_1(v_1),
self.quant_v_2(v_2),
], dim=1)
return k_features, v_features
model = MultiScaleCatModule()
# 关键:使用 qconfig_dict 为特定模块配置
# 2.3.6 版本可能支持通过模块名称前缀共享 observer
shared_qconfig = default_qat_8bit_fake_quant_qconfig()
qconfig_dict = {
"": default_qat_8bit_fake_quant_qconfig(),
"module_name": {
# K 路径共享
"quant_k_0": shared_qconfig,
"quant_k_1": shared_qconfig,
"quant_k_2": shared_qconfig,
# V 路径共享
"quant_v_0": shared_qconfig,
"quant_v_1": shared_qconfig,
"quant_v_2": shared_qconfig,
}
}
model = prepare_qat_fx(copy.deepcopy(model), qconfig_dict)
# calibration
model.eval()
set_fake_quantize(model, FakeQuantState.CALIBRATION)
with torch.no_grad():
for _ in range(10):
model(torch.rand(1, 128, 64, 64))
# validation
model.eval()
set_fake_quantize(model, FakeQuantState.VALIDATION)
# 检查 scale 是否共享
for name in ['quant_k_0', 'quant_k_1', 'quant_k_2']:
module = getattr(model, name)
if hasattr(module, 'activation_post_process'):
obs = module.activation_post_process
if hasattr(obs, 'scale') and obs.scale is not None:
print(f"{name}: scale={obs.scale.item():.6f}")
model = convert_fx(model)
方案 2:使用 PTQ 替代 QAT(更简单)
如果 QAT 太复杂,PTQ 的 YAML 配置可能更灵活:
# quant_config.yaml
compiler_parameters:
optimization_level: 2
quantization_parameters:
default_qconfig: "default_8bit"
# 关键:针对 cat 前的 layer 配置共享 observer
layer_parameters:
- layer_name: "quant_k_0"
share_observer: true
- layer_name: "quant_k_1"
share_observer: true
- layer_name: "quant_k_2"
share_observer: true
# 使用 hb_mapper 进行 PTQ
hb_mapper \
--march bernoulli2 \
--model-type onnx \
--input-model your_model.onnx \
--calibration-dataset ./calib_data \
--quant-config quant_config.yaml \
--output-path ./output
方案 3:联系 FAE 获取 2.3.6 版本特定的 API(最可靠)
正如 Marcelo6151 所说,QAT 深度优化需要 FAE 支持。2.3.6 是新版本,API 可能有变化。
建议提供以下信息给 FAE:
| 信息项 | 内容 |
|---|---|
| 工具链版本 | horizon-plugin-pytorch 2.3.6+cu116.torch1130 |
| 问题描述 | torch.cat 产生大量 hz_cat_1_rescale_1 算子 |
| 错误信息 | AssertionError: input scale must be the same as op's |
| 模型结构 | 多尺度特征 cat 操作 |
| 已尝试方案 | 共享 QuantStub、手动绑定 observer 等 |
方案 4:接受 rescale 开销,优化模型结构
如果上述方案都不可行,减少 rescale 的性能影响:
# 减少 cat 次数
# 不推荐
x = torch.cat([a, b], dim=1)
x = torch.cat([x, c], dim=1)
x = torch.cat([x, d], dim=1)
# 推荐:一次性 cat
x = torch.cat([a, b, c, d], dim=1)
# 或者用 add 替代(如果业务逻辑允许)
x = a + b + c + d # elementwise add 不需要 rescale
诊断代码
请运行以下代码,告诉我输出,方便进一步定位:
import horizon_plugin_pytorch
print(f"Plugin version: {horizon_plugin_pytorch.__version__}")
# 检查 QuantStub 的属性
from horizon_plugin_pytorch.quantization import QuantStub
q = QuantStub()
print(f"QuantStub attributes: {[x for x in dir(q) if not x.startswith('_')]}")
# 检查 prepare_qat_fx 的签名
from horizon_plugin_pytorch.quantization import prepare_qat_fx
import inspect
print(f"prepare_qat_fx signature: {inspect.signature(prepare_qat_fx)}")
# 检查是否有 share_observer 相关 API
from horizon_plugin_pytorch.quantization import qconfig
print(f"qconfig module: {[x for x in dir(qconfig) if 'share' in x.lower()]}")
建议
- 先运行诊断代码,告诉我输出
- 如果仍无法解决,建议联系 FAE 获取 2.3.6 版本特定的支持
- 参考社区资源:
有诊断结果随时同步 ![]()