PyTorch FX图模式量化感知训练(QAT)如何冻结BN统计量
结论
torch.nn.intrinsic.qat.freeze_bn_stats 在FX Graph模式的QAT流程中完全可用,你可以直接调用 model_prepared_fx.apply(torch.nn.intrinsic.qat.freeze_bn_stats) 实现冻结BN统计量的目标,不需要额外替换其他机制。
原理说明
quantize_fx.prepare_qat_fx 接口在处理模型时,会自动完成算子融合、插入量化观察者、替换普通算子为QAT版本算子的操作,最终输出的model_prepared_fx中所有BN相关模块(包括单独的QAT BN、Conv+BN、Conv+BN+ReLU等融合QAT模块),和eager模式QAT流程生成的对应模块类型完全一致,均继承自torch.nn.intrinsic.qat下的标准QAT模块类,因此freeze_bn_stats的适配逻辑可以直接生效。
完整FX QAT训练循环示例
你可以参考如下代码补全训练流程,逻辑和eager模式的官方示例完全对齐:
import copy import torch from torch.ao.quantization import quantize_fx # 替换为你自己的原始浮点模型、损失函数、优化器、数据加载器等组件 model_fp = your_custom_model criterion = your_loss_function optimizer = your_optimizer data_loader = your_train_dataloader data_loader_test = your_val_dataloader device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') eval_batch_size = your_eval_batch_size num_train_batches = 20 num_eval_batches = len(data_loader_test) num_epochs = 8 # FX QAT前置准备 model_to_quantize = copy.deepcopy(model_fp) qconfig_dict = {"": torch.quantization.get_default_qat_qconfig('qnnpack')} model_to_quantize.train() model_prepared = quantize_fx.prepare_qat_fx(model_to_quantize, qconfig_dict) # 训练循环 for nepoch in range(num_epochs): train_one_epoch(model_prepared, criterion, optimizer, data_loader, device, num_train_batches) # 冻结量化观察者参数,逻辑和eager模式一致 if nepoch > 3: model_prepared.apply(torch.ao.quantization.disable_observer) # 冻结BN运行时统计量,直接复用eager模式API即可生效 if nepoch > 2: model_prepared.apply(torch.nn.intrinsic.qat.freeze_bn_stats) # 验证精度 quantized_model = quantize_fx.convert_fx(model_prepared.eval(), inplace=False) quantized_model.eval() top1, top5 = evaluate(quantized_model, criterion, data_loader_test, neval_batches=num_eval_batches) print(f'Epoch {nepoch} :Evaluation accuracy on {num_eval_batches * eval_batch_size} images, {top1.avg:.2f}')
注意事项
- 请使用PyTorch 1.12及以上版本,低于该版本的FX QAT接口存在较多兼容性问题,可能出现模块类型不匹配的问题
- 如果你自定义了非标准的QAT融合模块,需要自行在模块中实现
freeze_bn_stats适配逻辑,默认的官方算子不需要额外处理
内容的提问来源于stack exchange,提问作者user8510613
相关产品推荐
相关产品推荐

