auto_channel_prune_search【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct产品支持情况产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√功能说明自动通道稀疏接口根据用户模型来计算各通道的稀疏敏感度影响精度以及稀疏收益影响性能然后搜索策略依据该输入来搜索最优的逐层通道稀疏率以平衡精度和性能。最终输出一个配置文件。函数原型auto_channel_prune_search(model, config, input_data, output_cfg, sensitivity, search_alg)参数说明参数名输入/输出说明model输入含义待稀疏的PyTorch模型。数据类型torch.nn.Moduleconfig输入含义自动通道稀疏配置文件路径。基于basic_info.proto文件中的AutoChannelPruneConfig生成的简易配置文件*.proto文件所在路径为AMCT安装目录/amct_pytorch/proto/。*.proto文件参数解释以及生成的自动通道稀疏搜索配置文件样例请参见自动通道稀疏搜索简易配置文件。数据类型stringinput_data输入含义用户提供获取输入数据含label。数据类型list[data,label]列表元素数据类型为torch.tensor。output_cfg输入含义输出的最终的通道稀疏配置文件路径。数据类型stringsensitivity输入含义敏感度计算方法。数据类型string或SensitivityBase的子类string为AMCT已有的方法目前可选为TaylorLossSensitivitySensitivityBase的子类实例化可由用户来继承定义。search_alg输入含义待稀疏的通道搜索方法。数据类型string或SearchChannelBase的子类string为AMCT已有的方法目前可选为GreedySearchSearchChannelBase的子类实例化可由用户来继承定义。返回值说明无调用示例import amct_pytorch as amct #构造输入数据input_data input_data torch.randn(input_shape) model.eval() output model.forward(input_data) labels torch.randn(output.size()) data [input_data,labels] amct.auto_channel_prune_search( modelmodel, config./tmp/sample.cfg, input_datadata, output_cfg./tmp/output.cfg, sensitivityTaylorLossSensitivity, search_algGreedySearch)落盘文件说明保存的自动通道稀疏配置文件需要传给通道稀疏接口完成后续的业务。【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考