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_model、save_compressed_retrain_model的完整配合流程。
一、产品支持情况
restore_compressed_retrain_model属于 AMCT 的静态组合压缩训练接口,其能力随目标硬件平台的不同有所差异。特性中标记为 "x" 的产品,调用接口本身不会报错,但无法获取对应的性能收益:
| 产品 | 量化感知训练 | 通道稀疏 | 4选2结构化稀疏 |
|---|---|---|---|
| Ascend 950PR / Ascend 950DT | INT8 量化 √;INT4 量化 x | √ | x |
| 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.Module,config_defination/record_file/pth_file必须是str,state_dict_name为str或None),随后通过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: M4N2,update_freq默认 0); - 全局/差异化配置:
skip_layers、skip_layer_types、quant_skip_layers、quant_skip_types、regular_prune_skip_layers、regular_prune_skip_types,以及按层/按层类型重写的override_layer_configs、override_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 中,主要步骤为:
- 前置处理:
ModuleHelper(model).check_amct_op()校验模型;尝试深拷贝模型;record_file、pth_file转为绝对路径;通过SingletonScaleOffsetRecord().reset_singleton(record_file)重置单例记录器以读取既有 record; - 恢复压缩结构:调用内部函数
_modify_original_to_compressed_model(model, input_data, config_defination, record_file, "restore"),与create_compressed_retrain_model共用同一套图变换逻辑,仅以prune_call_mode区分创建/恢复分支; - 加载权重:调用
load_pth_file(model, pth_file, state_dict_name)将 checkpoint 权重加载进重建后的模型; - 返回:返回修改后的
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 文件中已记录的稀疏关系(而非重新计算),这保证了恢复后的模型结构与训练中断前的压缩结构完全一致。
七、配套工作流与落盘产物
完整的静态组合压缩训练流程建议按以下顺序组织:
- 创建:
amct.create_compressed_retrain_model(model, input_data, config_defination, record_file)生成压缩训练模型,record 文件记录稀疏(若配置了稀疏)与量化因子; - 训练:对返回模型进行量化感知训练,定期
torch.save保存 checkpoint(键名记为state_dict); - 恢复:训练中断或结束后,用
amct.restore_compressed_retrain_model(model, input_data, config_defination, record_file, pth_file, 'state_dict')重建并加载权重,可继续训练或直接用于导出; - 导出:
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_defination与record_file必须与创建阶段一致;state_dict_name需与保存 checkpoint 时的键名对应。其内部与创建接口共享同一套图变换管线,确保了恢复结构的确定性,这也是断点续训与训练后重建能够稳定复现的前提。
【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考