auto_channel_prune_search
产品支持情况
| 产品 | 是否支持 |
|---|---|
| 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)
参数说明
|
基于basic_info.proto文件中的AutoChannelPruneConfig生成的简易配置文件,*.proto文件所在路径为:AMCT安装目录/amct_pytorch/proto/。 *.proto文件参数解释以及生成的自动通道稀疏搜索配置文件样例请参见自动通道稀疏搜索简易配置文件。 |
||
|
数据类型:string或SensitivityBase的子类,string为AMCT已有的方法,目前可选为'TaylorLossSensitivity';SensitivityBase的子类实例化,可由用户来继承定义。 |
||
|
数据类型:string或SearchChannelBase的子类,string为AMCT已有的方法,目前可选为'GreedySearch';SearchChannelBase的子类实例化,可由用户来继承定义。 |
返回值说明
无
调用示例
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(
model=model,
config='./tmp/sample.cfg',
input_data=data,
output_cfg='./tmp/output.cfg',
sensitivity='TaylorLossSensitivity',
search_alg='GreedySearch')
落盘文件说明:
保存的自动通道稀疏配置文件,需要传给通道稀疏接口完成后续的业务。