AMCT 组合压缩训练恢复:restore_compressed_retrain_model 接口详解
2026/9/18 22:07:03 网站建设 项目流程

AMCT 组合压缩训练恢复:restore_compressed_retrain_model 接口详解

【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct

本文基于 CANN AMCT 开源仓库(amct_pytorch)讲解静态组合压缩训练的恢复接口restore_compressed_retrain_model。该接口用于在「稀疏 + 量化」组合压缩的训练阶段中断或训练完成后,依据训练过程生成的 record 记录文件与 checkpoint 权重,重新构建可用于继续训练或导出的压缩训练模型。读完本文,你将掌握该接口的参数含义、配置准备、调用示例,以及其内部与create_compressed_retrain_modelsave_compressed_retrain_model的完整配合流程。

一、产品支持情况

restore_compressed_retrain_model属于 AMCT 的静态组合压缩训练接口,其能力随目标硬件平台的不同有所差异。特性中标记为 "x" 的产品,调用接口本身不会报错,但无法获取对应的性能收益:

产品量化感知训练通道稀疏4选2结构化稀疏
Ascend 950PR / Ascend 950DTINT8 量化 √;INT4 量化 xx
Atlas A3 训练系列产品 / Atlas A3 推理系列产品INT8 量化 √;INT4 量化 x
Atlas A2 训练系列产品 / Atlas A2 推理系列产品INT8 量化 √;INT4 量化 x

注意:当前版本量化感知训练仅支持 INT8 量化;4选2结构化稀疏在 Ascend 950PR/Ascend 950DT 上因硬件约束不受支持,相关配置项的详细说明可参见 量化感知训练简易配置文件。

二、功能说明:组合压缩训练的恢复入口

静态组合压缩训练的整体思路是「先稀疏、后量化」:先将原始模型按组合压缩配置执行通道稀疏(或4选2结构化稀疏),再插入量化相关的算子(数据和权重的量化感知训练层以及 searchN 层),随后在训练过程中保存 checkpoint 权重。

restore_compressed_retrain_model是这一流程中的恢复(restore)接口:它将传入的待压缩模型,按照给定的组合压缩配置文件config_defination和训练期间记录的 record 记录文件(含稀疏与量化因子),重新执行「先稀疏后量化」的图变换,并加载训练过程中保存的 checkpoint 权重参数,最终返回修改后的torch.nn.Module模型。

它与同一套流程中的另外两个接口配套使用,构成完整的生命周期:

  • create_compressed_retrain_model:首次创建压缩训练模型并生成 record 文件;
  • restore_compressed_retrain_model(本文):基于 record 文件恢复压缩结构并加载权重,用于断点续训或训练后重建;
  • save_compressed_retrain_model:将恢复后的模型导出为 deploy/fake quant 的 ONNX 文件。

从源码结构看,三个接口都定义在 prune_interface.py 中,并在 amct_pytorch 包入口 中统一导出为amct_pytorch的公共 API。

三、函数原型与参数说明

compressed_retrain_model = restore_compressed_retrain_model(model, input_data, config_defination, record_file, pth_file, state_dict_name=None)

3.1 参数详解

参数名输入/输出说明
model输入含义:PyTorch 的 model。数据类型:torch.nn.Module
input_data输入含义:模型的输入数据。一个torch.tensor会被等价为tuple(torch.tensor)。数据类型:tuple
config_defination输入含义:静态组合压缩简易配置文件。基于retrain_config_pytorch.proto文件生成的简易配置文件compressed.cfg.proto文件所在路径为:AMCT安装目录/amct_pytorch/proto/(仓库内对应 retrain_config_pytorch.proto)。参数解释及配置样例请参见 量化感知训练简易配置文件。数据类型:string
record_file输入含义:已经记录稀疏和量化因子的文件(由create_compressed_retrain_model生成)。数据类型:string
pth_file输入含义:训练过程中保存的权重文件(checkpoint)。数据类型:string
state_dict_name输入含义:权重文件中权重对应的键值。默认值:None。数据类型:string

从源码实现看,接口在进入核心逻辑前会通过@check_params装饰器完成类型校验(model必须是torch.nn.Moduleconfig_defination/record_file/pth_file必须是strstate_dict_namestrNone),随后通过ModuleHelper(model).check_amct_op()检查模型中是否已包含 AMCT 自定义算子,并尝试对模型做深拷贝,避免修改原始模型(见 prune_interface.py)。

3.2 返回值说明

返回根据record_file中的稀疏关系进行稀疏后、且插入量化相关层、并已加载权重文件的torch.nn.Module静态组合压缩训练模型。

3.3 约束说明

组合压缩配置文件至少存在一个配置:稀疏配置或者量化配置。

四、配置准备:组合压缩简易配置文件

config_defination指向的组合压缩简易配置文件基于retrain_config_pytorch.proto生成,语法与量化感知训练/稀疏简易配置同源(同一 proto 可配置出量化、稀疏、组合压缩三种场景),核心配置项包括:

  • 量化侧retrain_data_quant_config(数据量化,ULQ 算法,dst_type默认 INT8,支持clip_max_min初始上下限、fixed_min等)与retrain_weight_quant_config(权重量化,ARQ/ULQ 算法,支持channel_wise);
  • 稀疏侧prune_config下的filter_pruner(通道稀疏,balanced_l2_norm_filter_prune算法,prune_ratio稀疏率,推荐 0.2,ascend_optimized昇腾亲和优化建议为 true)或n_out_of_m_pruner(4选2结构化稀疏,l1_selective_prune算法,n_out_of_m_type: M4N2update_freq默认 0);
  • 全局/差异化配置skip_layersskip_layer_typesquant_skip_layersquant_skip_typesregular_prune_skip_layersregular_prune_skip_types,以及按层/按层类型重写的override_layer_configsoverride_layer_types。参数优先级为:override_layer_configs>override_layer_types> 全局量化/稀疏配置。

组合压缩(通道稀疏 + INT8 量化)简易配置文件compressed1.cfg示例(完整参数表与更多样例见 量化感知训练简易配置文件):

prune_config : { filter_pruner : { balanced_l2_norm_filter_prune : { prune_ratio : 0.3 ascend_optimized: True } } } # skip_layers: "skip_layers_name_0" skip_layer_types: "Optype" quant_skip_layers: "Opname" quant_skip_types: "Optype" retrain_weight_quant_config: { arq_retrain: { channel_wise: true dst_type: INT8 } } override_layer_types : { layer_type: "Optype" retrain_weight_quant_config: { arq_retrain: { channel_wise: false dst_type: INT8 } } retrain_data_quant_config : { ulq_quantize : { clip_max_min : { clip_max : 6.0 clip_min : -6.0 } } } prune_config : { filter_pruner : { balanced_l2_norm_filter_prune : { prune_ratio : 0.5 ascend_optimized: True } } } }

五、调用示例

以下调用示例完整演示了「建立模型 → 保存权重 → 恢复压缩训练模型」的流程(与仓库测试用例的用法一致):

import amct_pytorch as amct # 建立待进行组合压缩的网络图结构 model = build_model() input_data = tuple(torch.randn(input_shape)) save_pth_path = /your/path/to/save/tmp.pth record_file = os.path.join(TMP, 'compressed_record.txt') config_defination = './compressed_cfg.cfg' torch.save({'state_dict': model.state_dict()}, save_pth_path) compressed_retrain_model = amct.restore_compressed_retrain_model( model, input_data, config_defination, record_file, save_pth_path, 'state_dict')

示例中state_dict_name传入'state_dict',与torch.save({'state_dict': model.state_dict()}, ...)保存的键名一一对应。仓库测试用例 test_prune_interface.py 展示了标准的实战组合:先create_compressed_retrain_model生成压缩模型并做一次推理(使量化参数完成初始化),再用torch.save保存其state_dict,随后调用restore_compressed_retrain_model基于原模型、同一record_file与配置文件重建并加载权重,最后用save_compressed_retrain_model导出 ONNX(生成*_deploy_model.onnx*_fake_quant_model.onnx两类文件)。

5.1 使用要点

  • input_data仅用于编译模型图结构(Parser.export_onnx与图解析),可使用随机数据;
  • 恢复时传入的config_defination应与首次创建时保持一致,record_file必须是create_compressed_retrain_model生成的记录文件;
  • 实际训练中,应保存的是压缩后模型(即create_compressed_retrain_model的返回值)的state_dict,恢复接口负责把权重正确映射回重建出的压缩结构。

六、内部实现:restore 流程的源码级拆解

restore_compressed_retrain_model的核心逻辑在 prune_interface.py 中,主要步骤为:

  1. 前置处理ModuleHelper(model).check_amct_op()校验模型;尝试深拷贝模型;record_filepth_file转为绝对路径;通过SingletonScaleOffsetRecord().reset_singleton(record_file)重置单例记录器以读取既有 record;
  2. 恢复压缩结构:调用内部函数_modify_original_to_compressed_model(model, input_data, config_defination, record_file, "restore"),与create_compressed_retrain_model共用同一套图变换逻辑,仅以prune_call_mode区分创建/恢复分支;
  3. 加载权重:调用load_pth_file(model, pth_file, state_dict_name)将 checkpoint 权重加载进重建后的模型;
  4. 返回:返回修改后的torch.nn.Module

其中_modify_original_to_compressed_model(见 prune_interface.py)的详细流程为:

步骤1 解析:Parser.export_onnx 导出 ONNX 并解析为内部图,RetrainConfig.init 解析组合压缩配置(enable_retrain=True, enable_prune=True) 步骤2 通道稀疏:若 enable_prune 且为 restore 模式, prune_helper.restore_prune_model() 恢复 filter 稀疏结构, restore_selective_prune_record() 恢复记录中的稀疏关系 步骤3 选择稀疏:若 enable_prune,_modify_original_model_to_prune 插入稀疏训练相关算子 步骤4 量化插入:若 enable_retrain,_modify_original_model_to_quant 插入数据和权重的 量化感知训练层以及 searchN 层

可见恢复流程与创建流程共用同一套「稀疏 + 量化」图变换骨架,差异仅在于稀疏部分读取的是 record 文件中已记录的稀疏关系(而非重新计算),这保证了恢复后的模型结构与训练中断前的压缩结构完全一致。

七、配套工作流与落盘产物

完整的静态组合压缩训练流程建议按以下顺序组织:

  1. 创建amct.create_compressed_retrain_model(model, input_data, config_defination, record_file)生成压缩训练模型,record 文件记录稀疏(若配置了稀疏)与量化因子;
  2. 训练:对返回模型进行量化感知训练,定期torch.save保存 checkpoint(键名记为state_dict);
  3. 恢复:训练中断或结束后,用amct.restore_compressed_retrain_model(model, input_data, config_defination, record_file, pth_file, 'state_dict')重建并加载权重,可继续训练或直接用于导出;
  4. 导出amct.save_compressed_retrain_model(model, record_file, save_path, input_data)输出 deploy 与 fake quant 两类 ONNX 文件;若只有稀疏配置(仅剪枝场景),两类文件内容相同。

八、总结

restore_compressed_retrain_model是 AMCT 静态组合压缩训练闭环中的关键恢复入口,它把「稀疏记录 + 量化因子 + checkpoint 权重」三者重新组织为一个可继续训练、可导出部署的torch.nn.Module。使用时需注意三点:配置文件必须至少包含稀疏或量化之一;恢复所用config_definationrecord_file必须与创建阶段一致;state_dict_name需与保存 checkpoint 时的键名对应。其内部与创建接口共享同一套图变换管线,确保了恢复结构的确定性,这也是断点续训与训练后重建能够稳定复现的前提。

【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询