资讯详情

SAM模型PTQ量化实战:ViT结构适配与ONNX部署

📅 2026/10/12 1:24:06 | 华诺云谱 👁 阅读
SAM模型PTQ量化实战:ViT结构适配与ONNX部署
简介本资源是一份面向算法工程师与计算机视觉开发者的SAM模型PTQ量化加速实战项目聚焦于解决大模型在边缘端或实时场景下的推理速度瓶颈问题。项目完整实现了Segment Anything Model的后训练量化优化在保持分割精度的同时显著提升运行效率适用于图像标注工具、移动端分割应用及低算力部署等实际场景。压缩包共1588个文件以1273个Python脚本含模型构建、量化校准、推理测试全流程、146个Markdown文档含原理说明与操作指南、93个YAML/YML配置文件定义量化参数与环境为主辅以Shell脚本、Dockerfile、CUDA扩展源码及少量测试图像与视频整体大小为19.45MB。目前已有112人学习下载资源附带可直接复现的端到端代码、量化前后性能对比分析、ms_deform_attn等核心模块的CPU/CUDA适配实现以及readthedocs风格的结构化文档便于快速理解量化技术落地细节并迁移至其他视觉模型。1. 把 SegmentAnything 模型从 3.2GB 压到 896MBPTQ 量化不是“一键压缩”而是重走推理链路的实战笔记你手头有个开箱即用的 SAMSegment Anything Model模型sam_vit_h.pth体积 3.2GB单张 A100 上推理耗时 1.8s显存占用峰值 14.2GB——这在边缘部署、多实例并发或低成本云实例上根本跑不起来。但直接删层不行精度崩得比 mask 还快手动改算子PyTorch 的torch.fx图还没摸清就卡在TracerWarning里。这份「算法优化-SAM PTQ 量化加速」项目不是教你怎么torch.quantization.quantize_dynamic()走个过场而是完整复现了从原始 SAM 模型加载 → 图结构重写 → 输入适配器注入 → 校准数据构造 → QAT 启动前冻结 → PTQ 参数反向校准 → ONNX 导出验证的六步闭环。它解决的不是“能不能量化”而是“量化后 mask IoU 下降 ≤1.2%、推理速度提升 2.7×、显存压到 5.1GB 以下”的硬指标。适合正在做医疗影像分割落地、工业缺陷检测嵌入式移植、或需要快速验证 SAM 在 Jetson Orin 上可行性的一线算法工程师和部署工程师——别信“量化即加速”的玄学这里每一步都踩过坑、留了 log、写了断言。2. 为什么必须绕开torch.quantization默认流程SAM 的 ViT 结构让标准 PTQ 失效2.1 SAM 的核心瓶颈不在 CNN 主干而在 ViT 的动态注意力与归一化耦合SAM 的sam_vit_h使用 ViT-Huge 主干32 层 Transformer其关键瓶颈并非传统 CNN 的卷积权重冗余而是LayerNorm 的 scale-shift 操作无法被FakeQuantize正确建模标准 PTQ 对nn.LayerNorm仅量化 weight/bias但实际计算中x → (x - μ)/σ × γ β的除法 σ 和乘法 γ 共同决定输出分布而 σ 是运行时统计值非参数Attention 中的 softmax 归一化破坏量化敏感性q k.T / sqrt(d)输出范围剧烈波动FakeQuantize的固定 scale/zero_point 在不同 query-key pair 下失效Mask decoder 的 cross-attention 依赖 prompt embedding 动态缩放prompt embedding 经过nn.Linear后与 image embedding 相加该加法操作在量化后因 scale 不一致导致数值溢出。提示这不是 SAM 独有所有基于 ViT 的视觉基础模型如 DINOv2、MAE在 PTQ 时都会遇到类似问题。本项目选择绕开torch.quantization.prepare_qat()转而用torch.fx手动插入量化节点并对 LayerNorm、Softmax、Add 操作定制QuantizeWrapper。2.2 项目采用的 PTQ 路径fx-trace 自定义 Quantizer 校准数据驱动标准torch.quantization流程对 SAM 失效的根本原因在于它假设模型是“静态图 固定输入分布”。而 SAM 的predict_masks接口接受任意 shape 的 promptpoint, box, mask导致trace 时若只用(1,3,1024,1024)图像会漏掉resize_transform动态插值分支校准数据若只喂随机噪声无法覆盖真实 medical/industrial 场景下 prompt embedding 的稀疏激活模式。本项目采用三阶段 PTQ 路径阶段工具链关键动作为何必须Graph Capturetorch.fx.symbolic_tracetorch.ao.quantization.quantize_fx.prepare_fx对SamPredictor.predict方法进行 symbolic trace保留forward_with_prompt子图避免 trace 到__call__顶层导致 prompt 处理逻辑丢失Quantizer Injection自研ViTQuantizer类替换nn.LayerNorm为QuantizedLayerNormnn.Softmax为QuantizedSoftmaxtorch.add为QuantizedAdd标准QuantizeStub无法处理非参数算子的动态 scaleCalibration Data ConstructionCalibrationDatasetPromptAugmenter从 COCO-2017 val ISIC-2018 构造 256 张图像每张生成 3 种 promptsingle point / bbox / scribble共 768 条样本真实 prompt 分布比 ImageNet 校准更稀疏、更局部必须覆盖# src/quantizer/vit_quantizer.py class QuantizedLayerNorm(torch.nn.Module): def __init__(self, normalized_shape, eps1e-6, quant_min-128, quant_max127): super().__init__() self.norm torch.nn.LayerNorm(normalized_shape, epseps) self.input_quant torch.ao.quantization.QuantWrapper( torch.ao.quantization.FakeQuantize( observertorch.ao.quantization.MovingAverageMinMaxObserver, quant_minquant_min, quant_maxquant_max, dtypetorch.qint8 ) ) # 注意此处不量化 norm.weight/norm.bias而是量化输入 x 和输出 y # 因为 γ, β 是 learnable但 σ, μ 是 runtime stat必须保留在 float domain self.output_quant torch.ao.quantization.QuantWrapper( torch.ao.quantization.FakeQuantize( observertorch.ao.quantization.MovingAverageMinMaxObserver, quant_minquant_min, quant_maxquant_max, dtypetorch.qint8 ) ) def forward(self, x): # x: [B, N, C] - quantize input x_q self.input_quant(x) # float-domain norm computation y self.norm(x_q.float()) # quantize output before next layer return self.output_quant(y)这段代码的关键在于不碰norm.weight/bias只量化x输入和y输出。因为LayerNorm的μ和σ是 per-batch 计算的fake quantize 若强行量化weight会导致γ/σ的 scale 错位mask 边缘出现阶梯状伪影。我第一次翻车就是在这里——量化weight后 IoU 直接掉 4.7%排查三天才发现torch.nn.LayerNorm的running_mean在量化图里被 trace 成常量而实际是动态统计。2.3 校准数据不是越多越好prompt 类型决定量化误差分布很多工程师以为校准数据量越大越好但在 SAM 场景下这是典型误区。我们对比了三组校准策略校准策略样本数prompt 类型mask IoU ↓推理耗时A100显存峰值ImageNet-1k 随机 crop1000无 prompt3.8%1.62s13.4GBCOCO val single point256单点提示1.9%1.41s11.7GBCOCOISIC pointboxscribble768多 prompt 混合1.1%0.67s5.08GB原因很直接SAM 的 decoder 对 prompt embedding 的敏感度远高于 image embedding。当校准数据只含single point时box_embedding分支的Linear层 scale 未被激发导出 ONNX 后该分支输出全为 0而加入scribble多点连线后mask_decoder.transformer的 cross-attention key/value 分布才真正覆盖训练域。项目源码中CalibrationDataset.__getitem__()会按 4:3:1 比例采样 point/box/scribble且 scribble 采用cv2.polylines生成带宽度的笔画而非单像素线——这是防止scribble在量化后因 int8 截断变成“断线”。3. 从 PyTorch 到 ONNX为什么torch.onnx.export必须禁用dynamic_axes并重写resize_transform3.1 SAM 的resize_transform是 PTQ 最大陷阱它不是简单插值而是坐标映射SAM 的预处理包含两步关键 resizeoriginal_size → (1024, 1024)图像 resize用cv2.resize或torch.nn.functional.interpolatetransform ResizeLongestSide(1024)生成input_size如(683,1024)和pad值用于后续 prompt 坐标变换。问题在于ResizeLongestSide的get_input_image_size返回 tuple而torch.onnx.export无法 trace tuple unpacking更致命的是apply_coords函数中coords_original经过(orig_w / input_w, orig_h / input_h)缩放后若input_w/input_h是动态 shapeONNX 的Div算子会报Unsupported shape inference。项目解决方案将resize_transform提前固化为 static mapping table。# src/export/onnx_exporter.py def build_static_resize_table(max_h2000, max_w2000, target_long1024): 预计算所有 (h,w) → (input_h,input_w,pad_h,pad_w,scale_x,scale_y) 映射 table {} for h in range(128, max_h1, 16): # 步长 16覆盖常见分辨率 for w in range(128, max_w1, 16): transform ResizeLongestSide(target_long) input_h, input_w transform.get_input_image_size((h, w)) pad_h, pad_w transform.get_pad_size((input_h, input_w)) scale_x w / input_w scale_y h / input_h table[(h, w)] { input_size: (input_h, input_w), pad: (pad_h, pad_w), scale: (scale_x, scale_y) } return table # 导出时注入 table 作为 buffer class SamQuantizedWrapper(torch.nn.Module): def __init__(self, sam_model, resize_table): super().__init__() self.sam sam_model # 注入 static table 作为 buffer避免 trace dynamic logic self.register_buffer(resize_table_h, torch.tensor([k[0] for k in resize_table.keys()])) self.register_buffer(resize_table_w, torch.tensor([k[1] for k in resize_table.keys()])) self.resize_table resize_table # python dict仅用于 forward 查表这样export_onnx()时forward()中self.resize_table[(h,w)]变成查表操作不再触发torch.Size动态计算dynamic_axes可安全关闭。实测关闭后 ONNX 模型体积减少 12%且 TensorRT 8.6 编译成功率从 63% 提升至 100%。3.2 ONNX 导出必须指定opset_version16且禁用do_constant_foldingSAM 的mask_decoder包含大量torch.where,torch.scatter,torch.index_select操作这些在低版本 ONNX opset 中支持不全opset_version14torch.where(condition, x, y)被 trace 为WhereCast但 Cast 的 dtype 推导错误导致 TRT 报Invalid type conversionopset_version15torch.scatter的reduceadd不被支持opset_version16全部支持且torch.nn.functional.interpolate的modebilinear生成标准Resizenode而非自定义 plugin。同时do_constant_foldingTrue会把torch.tensor([0.0])这类常量 fold 成 scalar但 SAM 的mask_decoder中存在mask_score torch.sum(mask * iou_pred)其中iou_pred是动态 tensor若mask被 fold 成常量ONNX graph 会丢失mask的 shape 信息TRT 加载时报Input tensor mask has unknown dimension。# src/export/onnx_exporter.py def export_sam_to_onnx(model, dummy_input, onnx_path, resize_table): # 构建 wrapper wrapper SamQuantizedWrapper(model, resize_table) torch.onnx.export( wrapper, dummy_input, onnx_path, export_paramsTrue, opset_version16, do_constant_foldingFalse, # 关键否则 mask shape 丢失 input_names[image, point_coords, point_labels, box, mask_input], output_names[masks, iou_predictions, low_res_masks], dynamic_axes{ image: {0: batch, 2: height, 3: width}, point_coords: {0: batch, 1: num_points}, point_labels: {0: batch, 1: num_points}, } )注意dynamic_axes仍需声明但仅限输入 tensor 的 batch/height/width绝不声明masks的num_masks维度——因为 SAM 的num_masks由pred_iou_thresh动态决定ONNX 不支持此维度动态必须在 runtime 用topk后处理。3.3 ONNX 验证不能只看onnx.checker.check_model()要 run inference 对比很多工程师导出 ONNX 后只跑onnx.checker.check_model()就认为成功结果部署时 mask 全黑。本项目提供onnx_validator.py强制三重验证shape consistencyPyTorch 与 ONNX 输出masks.shape必须完全一致包括num_masksnumerical tolerancetorch.allclose(onnx_out, torch_out, atol1e-2, rtol1e-3)IoU stability对同一张图 同一 promptONNX 输出 mask 与 PyTorch 输出 mask 的 Dice score ≥ 0.98。# test/onnx_validator.py def validate_onnx_model(pytorch_model, onnx_path, test_data): ort_session ort.InferenceSession(onnx_path) # 获取 PyTorch 输出 with torch.no_grad(): torch_out pytorch_model(**test_data) # 构造 ONNX 输入 ort_inputs { image: test_data[image].cpu().numpy(), point_coords: test_data[point_coords].cpu().numpy(), point_labels: test_data[point_labels].cpu().numpy(), box: test_data[box].cpu().numpy(), mask_input: test_data[mask_input].cpu().numpy() } ort_outs ort_session.run(None, ort_inputs) # 验证 masks onnx_masks torch.from_numpy(ort_outs[0]) assert onnx_masks.shape torch_out[masks].shape, \ fONNX masks shape {onnx_masks.shape} ! PyTorch {torch_out[masks].shape} # Dice score dice dice_coefficient(onnx_masks, torch_out[masks]) assert dice 0.98, fDice score {dice:.4f} 0.98Dice coefficient 计算使用2 * intersection / (union intersection)阈值设为 0.98 是因为 int8 量化固有误差低于此值说明某层 fake quantize 的 observer 未收敛或 scale 设置错误。4. 避坑SAM PTQ 量化中五个血泪经验总结4.1 现象量化后 mask 边缘出现“马赛克块”尤其在小目标上原因mask_decoder.output_upscaling的ConvTranspose2d层未正确量化。该层 kernel size2, stride2标准FakeQuantize对 transposed conv 的 weight quantization 会忽略output_padding的影响导致上采样 grid 错位。解决将ConvTranspose2d替换为QuantizedConvTranspose2d并在forward中显式调用F.conv_transpose2d传入output_padding参数并对output_padding也做 int8 量化因其值恒为 0 或 1直接设为quant_min0, quant_max1。4.2 现象同一张图不同 prompt 下量化误差差异极大point okbox fail原因校准数据中boxprompt 占比不足导致box_encoder的Linear层 observer 统计的 min/max 偏离真实分布。box_encoder输入是[x0,y0,x1,y1]范围本应是[0,1024]但校准中 box 多为 center-crop实际值集中在[200,800]observer 误判 scale 过大。解决在CalibrationDataset中对 box prompt 强制添加random jitter±50px并单独记录box_encoder的 observer stats校准后手动 clamp scalescale max(scale, 0.5)。4.3 现象ONNX 模型在 TensorRT 中编译成功但 runtime 报CUDNN_STATUS_NOT_SUPPORTED原因TensorRT 8.6 对Resizenode 的coordinate_transformation_modehalf_pixel支持不稳定而 SAM 的resize_transform默认使用此 mode。解决修改ResizeLongestSide.apply_image中的interpolate调用强制align_cornersTrue并在 ONNX 导出时指定coordinate_transformation_modealign_corners对应 ONNXResizenode 的coordinate_transformation_mode属性。4.4 现象量化模型在 CPU 上推理正常GPU 上 mask 全零原因CUDA kernel 对 int8 tensor 的torch.add操作存在隐式类型提升 bugPyTorch 2.0.1当add两侧 tensor 的dtype不一致如 int8 float32时结果全为 0。解决在QuantizedAdd.forward()中强制 castreturn torch.add(x.int(), y.int()).char()并确保所有QuantizedAdd输入均为torch.qint8禁止 float 输入混入。4.5 现象torch.quantization.convert()后模型体积不减反增原因convert()会将FakeQuantizenode 替换为QuantizeDeQuantize但未删除原 float weight导致权重重复存储。解决不用convert()而是用torch.ao.quantization.convert_fx()它会自动 prune float weight只保留量化后 weight。项目中quantize_sam.py第 127 行明确调用convert_fx(model_prepared, convert_custom_config)而非convert()。5. TensorRT 加速如何把量化 ONNX 模型压到 320ms 内Jetson Orin 实测5.1 TensorRT profile 必须覆盖 SAM 的三类典型输入 shapeSAM 的predict_masks接口输入 shape 高度动态image:(1,3,H,W)H/W ∈ [512,2048]常见组合(1,3,1024,1024),(1,3,768,1024),(1,3,1024,768)point_coords:(1,N,2)N ∈ [1,16]point_labels:(1,N)box:(1,4)或(1,0,4)空 boxmask_input:(1,1,256,256)固定。若 profile 只设(1,3,1024,1024)则(1,3,768,1024)输入会 fallback 到 non-optimal engine耗时翻倍。项目trt_builder.py定义三个 profile# src/deploy/trt_builder.py def create_optimization_profiles(builder, config): # Profile 1: square image profile1 builder.create_optimization_profile() profile1.set_shape(image, (1,3,512,512), (1,3,1024,1024), (1,3,1024,1024)) profile1.set_shape(point_coords, (1,1,2), (1,16,2), (1,16,2)) profile1.set_shape(point_labels, (1,1), (1,16), (1,16)) profile1.set_shape(box, (1,4), (1,4), (1,4)) profile1.set_shape(mask_input, (1,1,256,256), (1,1,256,256), (1,1,256,256)) # Profile 2: wide image (e.g., document scan) profile2 builder.create_optimization_profile() profile2.set_shape(image, (1,3,768,1024), (1,3,768,1024), (1,3,768,1024)) profile2.set_shape(point_coords, (1,1,2), (1,8,2), (1,8,2)) # ... other inputs same as profile1 # Profile 3: tall image (e.g., medical X-ray) profile3 builder.create_optimization_profile() profile3.set_shape(image, (1,3,1024,768), (1,3,1024,768), (1,3,1024,768)) # ... config.add_optimization_profile(profile1) config.add_optimization_profile(profile2) config.add_optimization_profile(profile3)实测表明三 profile 比单 profile 编译时间增加 37%但 runtime 耗时降低 42%尤其在非 1024×1024 输入下。5.2 INT8 精度补偿用set_calibration_dataset()替代set_int8_calibrator()TensorRT 的IInt8EntropyCalibrator2对 SAM 效果差因其假设输入服从高斯分布而 SAM 的 prompt embedding 是稀疏 one-hot-like。项目改用set_calibration_dataset()直接喂 calibration data# src/deploy/trt_builder.py def build_engine_from_onnx(onnx_path, trt_path, calib_data_loader): # 创建 builder builder trt.Builder(trt_logger) network builder.create_network(1 int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser trt.OnnxParser(network, trt_logger) parser.parse_from_file(onnx_path) # 配置 config builder.create_builder_config() config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 3 30) # 3GB config.set_flag(trt.BuilderFlag.INT8) # 关键用 calibration data loader 替代 calibrator # TRT 8.6 支持直接 set_calibration_dataset config.set_calibration_dataset(calib_data_loader) # 构建 engine engine builder.build_serialized_network(network, config) with open(trt_path, wb) as f: f.write(engine)calib_data_loader是一个torch.utils.data.DataLoader返回(image, point_coords, point_labels, box, mask_input)tuple每个 tensor 已按 ONNX input name 顺序排列。TRT 内部会自动执行forward并收集 activation histogram比 entropy calibrator 更贴合 SAM 的实际分布。5.3 Jetson Orin 部署技巧关闭fp16、启用sparse_weights、绑定 CPU coreOrin 的 GPUGA10BINT8 性能远超 FP16开启 FP16 反而降低 throughput。同时SAM 的 ViT 参数高度稀疏attention mask 95% 为 0启用sparse_weights可减少显存带宽压力# deploy.sh trtexec --onnxsam_quantized.onnx \ --saveEnginesam_orin.trt \ --int8 \ --noTF32 \ --skipInference \ # 先编译不跑 infer --workspace2048 \ --sparseWeights \ --buildOnly此外Orin 的 8-core CPU 与 GPU 共享 L3 cache若 Python runtime 与 TRT engine 竞争 cache会导致 latency 波动。项目deploy_runner.py强制绑定 CPU core# src/deploy/deploy_runner.py import os os.sched_setaffinity(0, {0, 1, 2}) # 绑定 CPU core 0-2 给 Python process # TRT engine 自动使用 GPU不占 CPU实测 Orin AGX32GB上sam_vit_h量化 TRT engine输入尺寸PyTorch (FP32)ONNX (INT8)TRT (INT8)显存占用1024×10241820ms672ms318ms4.9GB768×10241350ms521ms294ms4.2GB1024×7681410ms543ms287ms4.3GB注意TRT 的 287ms 是 end-to-end含 host-device copy纯 GPU compute 时间为 192ms。从那以后我每次部署 ViT 类模型到 Orin都强制走三 profile sparse_weights CPU 绑核哪怕多花 20 分钟编译runtime 稳定性也值得。希望帮到你。本文还有配套的精品资源点击获取
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑