
AMCTcreate_prune_retrain_model接口详解基于通道稀疏与 4 选 2 结构化稀疏构建稀疏后训练模型【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct导读本文聚焦 CANN AMCT昇腾 AI 处理器亲和的模型压缩工具为 PyTorch 模型提供的稀疏化后训练入口create_prune_retrain_model它接收一个已加载权重的torch.nn.Module图结构依据简易配置文件prune.cfg完成通道稀疏或4 选 2 结构化稀疏处理插入/替换相关算子生成记录稀疏信息的record_file并返回可直接用于稀疏后训练的模型。读完本文你将掌握该接口的完整参数语义、两种稀疏特性的适用层与约束、配置文件写法以及其底层在图解析、PruneHelper、选择性稀疏插入等环节的实现原理从而在昇腾环境上正确落地剪枝 重训练压缩流程。功能说明通道稀疏或 4 选 2 结构化稀疏接口两种稀疏特性每次只能使能一个将输入的待稀疏图结构按照给定的稀疏配置文件进行稀疏处理在传入的图结构中插入或替换相关算子生成记录稀疏信息的record_file返回修改后可用于稀疏后训练的torch.nn.Module模型。两种稀疏特性的能力差异如下稀疏特性粒度典型目标通道稀疏按 filter/通道整条移除压缩torch.nn.Linear输出维度与torch.nn.Conv2d输出通道4 选 2 结构化稀疏权重矩阵内每 4 个元素保留 2 个生成昇腾硬件友好的结构化稀疏权重M4N2需要特别说明的是create_prune_retrain_model生成的是稀疏后训练模型用于继续训练以恢复精度其下游还需要配套的restore_prune_retrain_model恢复稀疏后训练得到的权重与save_prune_retrain_model导出 ONNX等接口共同组成完整流程。产品支持情况产品是否支持Ascend 950PR / Ascend 950DT通道稀疏√4 选 2 结构化稀疏接口xAtlas A3 训练系列产品 / Atlas A3 推理系列产品通道稀疏√4 选 2 结构化稀疏接口√Atlas A2 训练系列产品 / Atlas A2 推理系列产品通道稀疏√4 选 2 结构化稀疏接口√注上述 4 选 2 结构化稀疏特性标记x的产品调用接口不会报错但是获取不到性能收益。函数原型prune_retrain_model create_prune_retrain_model(model, input_data, config_defination, record_file)该接口在仓库中的实际实现位于 prune_interface.py并通过amct_pytorch顶层导出调用方式为amct.create_prune_retrain_model(...)。参数说明参数名输入/输出说明model输入含义待进行稀疏的模型已加载权重。数据类型torch.nn.Moduleinput_data输入含义模型的输入数据。一个torch.tensor会被等价为tuple(torch.tensor)。数据类型tupleconfig_defination输入含义简易配置文件。基于retrain_config_pytorch.proto文件生成的简易配置文件prune.cfg*.proto 文件所在路径为AMCT安装目录/amct_pytorch/proto/。*.proto 文件参数解释以及生成的prune.cfg简易配置文件样例参见量化感知训练简易配置文件。数据类型stringrecord_file输入含义记录稀疏信息的文件路径及名称记录通道稀疏结点间的级联关系或记录 4 选 2 稀疏的节点。数据类型string返回值修改后可用于稀疏后训练的torch.nn.Module模型。通道稀疏支持的层及约束优化方式支持的层类型约束通道稀疏torch.nn.Linear全连接层复用层共用 weight 和 bias 参数不支持稀疏。通道稀疏torch.nn.Conv2d卷积层复用层共用 weight 和 bias 参数不支持稀疏depthwise 只能被动稀疏groups in_channels不能主动稀疏只支持 input data 的 shape 为(N, Cin, Hin, Win)。4 选 2 结构化稀疏支持的层及约束优化方式支持的层类型约束4 选 2 结构化稀疏torch.nn.Linear全连接层复用层共用 weight不支持稀疏。4 选 2 结构化稀疏torch.nn.Conv2d卷积层复用层共用 weight不支持稀疏只支持 input data 的 shape 为(N, Cin, Hin, Win)。4 选 2 结构化稀疏torch.nn.ConvTranspose2d反卷积层复用层共用 weight不支持稀疏只支持 input data 的 shape 为(N, Cin, Hin, Win)。可以看到两种稀疏特性均不支持复用层多个层共享同一份 weight/bias 参数——这是出于参数一致性的考虑共用参数无法独立地对单一层做结构变换。此外两种特性的 CNN 场景都限定输入为 4 维(N, Cin, Hin, Win)。简易配置文件prune.cfg的写法与参数解释config_defination指向的prune.cfg是基于retrain_config_pytorch.proto生成的简易配置文件。该 proto 文件在仓库中的位置为 retrain_config_pytorch.proto其稀疏相关结构可归纳如下message PruneConfig { oneof regular_prune_strategy { FilterPruner filter_pruner 1; // 通道稀疏 NOutOfMPruner n_out_of_m_pruner 2; // 4 选 2 结构化稀疏 } } message FilterPruner { oneof filter_pruner_algo { BalancedL2NormFilterPruner balanced_l2_norm_filter_prune 1; } } message BalancedL2NormFilterPruner { required float prune_ratio 1; // 通道稀疏比例 optional bool ascend_optimized 2 [default true]; // 昇腾亲和优化 } message NOutOfMPruner { oneof n_out_of_m_pruner_algo { L1SelectivePruner l1_selective_prune 1; } } message L1SelectivePruner { optional NOutOfMType n_out_of_m_type 1 [default M4N2]; // 4 选 2 optional uint32 update_freq 2 [default 0]; }通道稀疏配置样例仓库测试用例中使用的实际配置 model_001_prune.cfg 如下prune_config : { filter_pruner: { balanced_l2_norm_filter_prune: { prune_ratio: 0.5 } } }关键字段说明prune_config全局稀疏参数通道稀疏与 4 选 2 结构化稀疏在此互斥选择oneof结构保证每次只能使能一个。filter_pruner.balanced_l2_norm_filter_prune.prune_ratio通道稀疏比例取值建议在(0, 1)区间例如0.3表示移除约 30% 的通道/filter。该算法基于 L2 范数做平衡的滤波器筛选避免过度集中于某一段。ascend_optimized是否启用昇腾亲和优化默认true即尽量让剪枝后的通道数与昇腾硬件计算约束对齐。全局跳过参数regular_prune_skip_layers与regular_prune_skip_types可按层名/层类型跳过不做稀疏的层若同时配置了skip_layers/skip_layer_types全局参数则取两者并集。逐层覆盖可通过override_layer_configs按层名与override_layer_types按层类型为特定层单独指定prune_config参数优先级为override_layer_configs override_layer_types prune_config详见量化感知训练简易配置文件。4 选 2 结构化稀疏配置样例prune_config{ n_out_of_m_pruner { l1_selective_prune { n_out_of_m_type: M4N2 } } }关键字段说明n_out_of_m_pruner.l1_selective_prune.n_out_of_m_type结构化稀疏类型当前仅支持M4N2即每 4 个权重元素中保留 L1 范数最大的 2 个4 选 2。update_freq稀疏 mask 的更新频率默认 0。在稀疏后训练过程中可按该频率周期性重新计算保留的 2 个元素从而让稀疏位置随训练演化进一步提升精度恢复效果。完整的 proto 参数解释与更多配置样例含 override 逐层配置、4 选 2 配置等请参见仓库文档 qat_config.md。底层实现原理接口内部调用链从源码结构看create_prune_retrain_model的核心执行流程prune_interface.py如下模型检查与深拷贝通过ModuleHelper(model).check_amct_op()校验模型是否已含 AMCT 算子随后ModuleHelper.deep_copy(model)深拷贝模型避免破坏用户原始模型对象。record_file 初始化files_util.create_empty_file(record_file, check_existTrue)创建空的记录文件并通过SingletonScaleOffsetRecord().reset_singleton(record_file)重置单例记录器后续所有稀疏信息通道级联关系或 4 选 2 稀疏节点都会写入该文件。图解析Parser.export_onnx(model, input_data, model_onnx)将模型导出为 ONNX 中间表示再由Parser.parse_net_to_graph解析为内部图结构并graph.add_model(model)建立图与原始模型的映射。配置解析RetrainConfig.init(graph, config_defination, enable_retrainTrue, enable_pruneTrue)解析prune.cfg按override 层 override 类型 全局的优先级生效。通道稀疏PruneHelper(graph, input_data, record_file).create_prune_model()依据prune_ratio执行滤波器级通道稀疏同时将 producer-consumer 之间的级联关系写入record_file保证后续层同步收缩维度。选择性稀疏4 选 2create_selective_prune_record(graph)记录 4 选 2 稀疏节点来自custom_op/selective_prune/selective_prune.py。算子插入/替换_modify_original_model_to_prune(model, graph)通过GraphOptimizer与InsertRetrainPrunePass在模型中插入或替换稀疏相关算子最终返回修改后的模型。该实现与测试用例 test_prune_interface.py 中的行为完全对应测试在调用create_prune_retrain_model后直接断言剪枝后的维度例如test_prune_model_002断言new_model.layer1[0].out_channels 120、test_prune_model_003断言layer1.out_features 512且layer3.out_features 128同时校验剪枝前后输出 shape 一致ori_output.shape new_output.shape证明模型经过稀疏处理后结构已真实收缩、前向仍可正常执行。调用示例import amct_pytorch as amct # 建立待进行稀疏的网络图结构 model build_model() model.load_state_dict(torch.load(state_dict_path)) input_data tuple([torch.randn(input_shape)]) # 调用稀疏模型 API record_file os.path.join(TMP, scale_offset_record.txt) cfg_file ./prune_config.cfg prune_retrain_model amct.create_prune_retrain_model( model, input_data, cfg_file, record_file)使用要点model必须是已加载权重的模型这与下游restore_prune_retrain_model要求未加载权重形成对照。input_data用于编译/解析图结构只需 shape 与真实输入一致可以是随机数据。record_file路径可自定义接口会创建该文件请务必保留好它——后续restore_prune_retrain_model需要读取同一record_file来保证恢复出的模型与稀疏模型结构一致参见 restore_prune_retrain_model。与配套接口的完整流程create_prune_retrain_model是稀疏后训练流程的入口完整链路通常为create_prune_retrain_model生成稀疏后训练模型与record_file使用返回的模型进行稀疏后训练即 retrain保存训练得到的权重pth_filerestore_prune_retrain_model基于同一record_file与config_defination恢复稀疏结构并载入pth_file权重save_prune_retrain_model导出*_deploy_model.onnx与*_fake_quant_model.onnx供昇腾部署使用。上述完整闭环在 test_prune_interface.py 的test_mix_prune_retrain_model中有端到端验证创建稀疏模型 → 保存权重 → 恢复模型 → 导出 ONNX并校验恢复前后模型输出完全一致(new_output new_output2).all()为 True。总结create_prune_retrain_model是 AMCT 面向 PyTorch 的稀疏后训练核心入口通过一个prune.cfg即可在通道稀疏与 4 选 2 结构化稀疏之间二选一自动完成图解析、稀疏记录、算子插入与模型改造。理解其产品支持差异、层类型约束、配置优先级、record_file 作用四个关键点是正确使用该接口并跑通稀疏 → 重训练 → 恢复 → 导出全流程的前提仓库中的 prune_interface.py、retrain_config_pytorch.proto 以及 test_prune_interface.py 提供了最直接的源码级参考。【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考