AMCT PyTorch 蒸馏量化接口 create_distill_model 完全指南:将浮点模型改造为可蒸馏的量化压缩模型 📅 发布时间:2026/9/18 14:00:22 👁 浏览次数: AMCT PyTorch 蒸馏量化接口 create_distill_model 完全指南将浮点模型改造为可蒸馏的量化压缩模型【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct导读本文介绍 CANN AMCT昇腾 AI 处理器亲和的模型压缩工具PyTorch 场景下的蒸馏量化接口amct.create_distill_model。该接口是蒸馏Distillation量化训练链路中的核心环节它接收用户基于 create_distill_config 生成的蒸馏量化配置文件对已加载权重的浮点模型执行图结构解析与量化算子插入数据和权重的蒸馏量化层以及找 N的层返回一个可直接参与蒸馏训练的新torch.nn.Module模型。读完本文你将掌握create_distill_model的完整调用方式、参数约束、配置文件的生成与参数语义以及它和create_distill_config、distill、save_distill_model组成的蒸馏量化全流程。一、功能说明create_distill_model 在蒸馏量化中的定位create_distill_model是 AMCT PyTorch 提供的蒸馏量化接口。其核心职责是将输入的待量化压缩的图结构按照给定的蒸馏量化配置文件进行量化处理在传入的图结构中插入量化相关的算子数据和权重的蒸馏量化层以及找 N 的层返回修改后可用于蒸馏的torch.nn.Module模型。这里的找 N 的层指的是为量化查找合适的位宽/缩放参数N所引入的辅助结构蒸馏训练阶段会通过损失函数反向传播对这些量化参数进行学习更新。也就是说该接口做的是模型改造而不是模型训练——改造完成后需要调用蒸馏接口distill完成实际的蒸馏训练再通过save_distill_model导出部署模型。从源码实现看该接口定义在 amct_pytorch/classic/graph_based/amct_pytorch/distillation_interface.py并通过amct_pytorch包对外导出见 amct_pytorch/classic/graph_based/amct_pytorch/init.py因此用户可直接通过import amct_pytorch as amct后以amct.create_distill_model(...)的方式调用。二、产品支持情况产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√蒸馏量化能力在昇腾训练/推理系列产品上均有支持接口行为一致。三、函数原型compress_model create_distill_model(config_file, model, input_data)四、参数说明参数名输入/输出说明config_file输入含义用户生成的蒸馏量化配置文件用于指定模型 network 中量化层的配置情况和蒸馏结构。数据类型string。使用约束该接口输入的 config.json必须与 create_distill_config 接口输入的 config.json 一致。model输入含义待进行蒸馏量化的原始浮点模型已加载权重。数据类型torch.nn.Module。input_data输入含义模型的输入数据。一个torch.tensor会被等价为tuple(torch.tensor)。数据类型tuple。补充说明来自源码约束config_file在进入主流程前会经过 distill_helper.py 中DistillHelper.get_config_file的校验——要求路径合法且文件真实存在否则抛出OSError。model若为torch.nn.parallel.DistributedDataParallel包装模型接口会自动解包取其.module参与改造。input_data的 shape 应尽量与真实推理输入一致因为它不仅用于模型导出解析也会影响图结构中中间张量 shape 的确定。五、返回值说明返回修改后可用于蒸馏的 torch.nn.Module 模型即compress_model。该模型结构上等同于原浮点模型但其中的可蒸馏量化层如 Conv2d、Linear已被替换为蒸馏量化模块插入了数据蒸馏量化层、权重蒸馏量化层等算子模型参数中同时包含了原始权重与可学习的量化参数如激活的acts_clip_max/acts_clip_min、权重的wts_scales/wts_offsets这些参数会在后续distill训练中被分组更新。六、调用示例原文档给出的最小可用示例import amct_pytorch as amct # 建立待进行蒸馏量化的网络图结构 model build_model() model.load_state_dict(torch.load(state_dict_path)) input_data tuple([torch.randn(input_shape)]) # 生成压缩模型 compress_model amct.create_distill_model( config_json_file, model, input_data)其中config_json_file必须是通过 create_distill_config 生成的蒸馏量化配置文件路径推荐放在./configs/config.json。为了保证示例可直接运行一个完整的最小流程通常如下import torch import amct_pytorch as amct # 1. 建立已加载权重的浮点模型 model build_model() # 用户自建 torch.nn.Module model.load_state_dict(torch.load(state_dict_path)) model.eval() input_data tuple([torch.randn(input_shape)]) # 2. 生成蒸馏量化配置文件可选通过 config_defination 传入简易 cfg 约束生成 amct.create_distill_config( config_file./configs/config.json, modelmodel, input_datainput_data, config_definationNone) # 3. 生成压缩模型蒸馏量化后的学生模型 compress_model amct.create_distill_model( config_file./configs/config.json, modelmodel, input_datainput_data) # 4. 执行蒸馏训练详见本文第八节 # 5. 导出部署模型注意create_distill_config与create_distill_model两次调用必须使用同一个config.json以保证配置文件中的层名、蒸馏结构与实际模型图结构一一对应。七、前置流程蒸馏量化配置文件的生成create_distill_model的改造行为完全由蒸馏量化配置文件驱动因此理解配置文件是正确使用本接口的前提。配置文件由 create_distill_config 接口自动生成amct.create_distill_config( config_file./configs/config.json, model, input_data, config_defination./configs/distill.cfg)其工作原理见 distillation_interface.py为将模型导出为 ONNX 图结构Parser.export_onnx解析为图对象Parser.parse_net_to_graph后自动找出所有可蒸馏量化的层和可蒸馏量化的结构将量化配置与蒸馏结构写入config_file。当config_definationNone时使用默认配置否则基于distill_config_pytorch.proto生成的简易配置文件distill.cfg来约束配置内容proto 文件参数详解与 cfg 样例参见 蒸馏简易配置文件。生成的 JSON 蒸馏量化配置文件INT8 场景样例如下{ version: 1, batch_num: 1, group_size: 1, data_dump: false, distill_group: [ [ conv1, bn, relu ], [ conv2, bn2, relu2 ] ], conv1: { quant_enable: true, distill_data_config: { algo: ulq_quantize, dst_type: INT8 }, distill_weight_config: { algo: arq_distill, channel_wise: true, dst_type: INT8 } }, conv2: { quant_enable: true, distill_data_config: { algo: ulq_quantize, dst_type: INT8 }, distill_weight_config: { algo: arq_distill, channel_wise: true, dst_type: INT8 } } }各字段语义结合 distill_field.py 的解析与校验逻辑version配置文件版本号由Version字段管理用于兼容性校验。batch_num蒸馏 batch 数量用于 ifmr 积累数据计算量化因子见 蒸馏简易配置文件。group_size蒸馏 block 中最小蒸馏单元个数源码要求其必须为大于 0 的整数默认值为 1GroupSize类val 0时抛ValueError。data_dumpteacher 网络 block 输入输出 dump 开关布尔类型默认falseDataDump类。置为true时蒸馏阶段会先对 teacher 模型的 block 输入输出做 dump再加载 dump 数据参与蒸馏。distill_group蒸馏结构分组每个元素是一个层名列表指定按块蒸馏的起始层到结束层蒸馏结构中仅支持torch.nn.Module类型的算子。各层配置conv1、conv2等key 为层名quant_enable该层是否参与量化默认trueQuantEnable类若所有层均为false配置解析会直接报错没有层启用蒸馏。distill_data_config数据激活蒸馏量化配置algo目前仅支持ulq_quantizeULQ 数据量化算法算法说明见 algorithm_brief.md可选dst_typeINT4/INT8默认 INT8当前版本仅支持 INT8。distill_weight_config权重蒸馏量化配置algo仅支持arq_distill默认ARQ 权重量化算法或ulq_distillULQ 权重量化算法channel_wise表示是否做 channel wise 量化dst_type默认 INT8。校验约束同一层激活与权重的dst_type必须一致DistillRootConfig.check_layer_config_legal否则抛ValueError。八、接口底层实现剖析create_distill_model 做了什么从 distillation_interface.py 的源码可以看到create_distill_model的执行链路如下参数校验check_params装饰器校验config_file为 str、model为torch.nn.Module、input_data为torch.Tensor或 tuple。深拷贝模型ModuleHelper.deep_copy(model)拷贝一份模型所有改造都作用于副本避免污染用户原始模型若深拷贝失败如存在不可复制的资源则回退到原模型并给出告警日志。DDP 解包若模型是torch.nn.parallel.DistributedDataParallel取其.module作为改造对象。AMCT 算子检查ModuleHelper(model).check_amct_op()检查模型是否已包含 AMCT 插入的算子防止重复改造。图结构解析Parser.export_onnx(model, input_data, tmp_onnx)将模型导出为 ONNX 中间表示再Parser.parse_net_to_graph解析为图对象并绑定模型graph.add_model(model)供后续匹配层与蒸馏结构使用。解析蒸馏配置parse_distill_config(config_file, model)读取 JSON 配置文件将其中的蒸馏结构distill_group与各层量化配置解析为内部 dict。插入量化算子构造ModelOptimizer注册InsertQatPass(distill_config)后执行optimizer.do_optimizer(model, graph)完成对模型的量化改造。InsertQatPass见 optimizer/insert_qat_pass.py内部维护了一个替换表REPLACE_DICT {Conv2d: Conv2dQAT, Linear: LinearQAT}即对图结构中匹配到的Conv2d、Linear等可蒸馏层依据配置中该层的distill_data_config/distill_weight_config替换为对应的蒸馏量化模块Conv2dQAT、LinearQAT并在模块中嵌入数据和权重的蒸馏量化层。蒸馏结构的匹配遵循配置中的distill_group分组。九、蒸馏训练create_distill_model 的后续接力接口create_distill_model返回的compress_model不会自动完成量化参数的训练需要调用蒸馏接口distill接力完成同一定义于 distillation_interface.pyamct.distill( model, # teacher黄金模型原始浮点模型 compress_model, # student 模型create_distill_model 的返回值 config_file, # 与 create_distill_model 相同的蒸馏配置文件 train_loader, # torch.utils.data.DataLoader epochs1, # 蒸馏轮数 lr1e-3, # 学习率 sample_instanceNone, # 自定义样本处理实例 lossNone, # 默认 torch.nn.MSELoss optimizerNone, # 默认 AdamW按参数分组设置学习率 )该接口的核心行为结合 distill_helper.pyifmr 校准do_calibration按batch_num配置的批次运行学生模型前向完成 ifmr增量式量化因子初始化。分块蒸馏按配置中的distill_group逐组蒸馏每组内依次前向各蒸馏模块以 teacher 模型对应 block 的输出作为目标用损失函数默认MSELoss计算 loss 并反向传播。分组学习率gen_optimizer_per_group会将蒸馏模块的参数分为三组普通参数继承lr、激活量化参数acts_clip_max/acts_clip_min学习率 0.1、权重量化参数wts_scales/wts_offsets学习率 0.0001默认使用 AdamW 优化器。蒸馏完成后还可调用save_distill_model见 distillation_interface.py对蒸馏模型执行删除蒸馏模块、BN 融合、导出 fakequant ONNX 与 deploy ONNX等后处理生成可部署的量化模型。十、使用注意事项与约束汇总配置文件一致性create_distill_model与create_distill_config必须使用同一个config.json配置文件路径必须真实存在否则接口会直接报错。模型状态传入的model必须是已加载权重的浮点模型接口内部会深拷贝模型原模型对象不受影响可继续作为 teacher 模型参与蒸馏。输入数据input_data建议使用与真实推理一致的 shape 的随机数据即可示例中torch.randn(input_shape)一个 tensor 会被等价为单元素 tuple。支持范围当前版本蒸馏量化仅支持 INT8 位宽dst_type默认 INT8配置解析对 INT4 预留了接口但当前仅 INT8 可用数据量化算法仅支持ulq权重量化算法支持arq_distill与ulq_distill。蒸馏结构约束distill_group中的层必须为torch.nn.Module类型算子且层名必须能在模型图中匹配否则distill阶段会抛出layer xxx get module failed错误。产品适配当前能力适配 Ascend 950PR/Ascend 950DT、Atlas A3 与 Atlas A2 训练/推理系列产品接口在昇腾环境上执行依赖昇腾软件栈与 torch_npu 等运行环境。通过本文介绍开发者可以完整掌握create_distill_config生成配置 →create_distill_model改造模型 →distill蒸馏训练 →save_distill_model导出部署的蒸馏量化全链路并在 AMCT 仓库中amct_pytorch/classic/graph_based/amct_pytorch/对照源码进一步理解每个环节的实现细节。【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考