
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 950DTINT8 量化 √INT4 量化 x√xAtlas 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_nameNone)3.1 参数详解参数名输入/输出说明model输入含义PyTorch 的 model。数据类型torch.nn.Moduleinput_data输入含义模型的输入数据。一个torch.tensor会被等价为tuple(torch.tensor)。数据类型tupleconfig_defination输入含义静态组合压缩简易配置文件。基于retrain_config_pytorch.proto文件生成的简易配置文件compressed.cfg.proto文件所在路径为AMCT安装目录/amct_pytorch/proto/仓库内对应 retrain_config_pytorch.proto。参数解释及配置样例请参见 量化感知训练简易配置文件。数据类型stringrecord_file输入含义已经记录稀疏和量化因子的文件由create_compressed_retrain_model生成。数据类型stringpth_file输入含义训练过程中保存的权重文件checkpoint。数据类型stringstate_dict_name输入含义权重文件中权重对应的键值。默认值None。数据类型string从源码实现看接口在进入核心逻辑前会通过check_params装饰器完成类型校验model必须是torch.nn.Moduleconfig_defination/record_file/pth_file必须是strstate_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.2ascend_optimized昇腾亲和优化建议为 true或n_out_of_m_pruner4选2结构化稀疏l1_selective_prune算法n_out_of_m_type: M4N2update_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_configsoverride_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_retrainTrue, enable_pruneTrue 步骤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),仅供参考