如何缩小Android端PyTorch VGG模型的ptl文件体积?
问题与代码
用户代码如下:
import torch # import joblib from torch.utils.mobile_optimizer import optimize_for_mobile from torchvision.models.vgg import vgg16 import torch, torchvision.models # lb = joblib.load('lb.pkl') device = torch.device('cuda:0') #device = torch.device('cpu')#'cuda:0') torch.backends.cudnn.benchmark = True model = vgg16().to(device) # model = torchvision.models.vgg16() path = 'model-22222.pth' torch.save(model.state_dict(), path) # nothing else here model.load_state_dict(torch.load(path)) #model.load_state_dict(torch.load('./model-76-0.7754.pth')) scripted_module = torch.jit.script(model) optimized_scripted_module = optimize_for_mobile(scripted_module) optimized_scripted_module._save_for_lite_interpreter("model-76-0.7754.ptl")
使用optimize_for_mobile后生成的ptl文件约527M,Android设备上体积过大,需缩小体积。
解决方案
1. 模型量化(最有效压缩方式)
量化将32位浮点参数转为8位整数,体积可压缩至原1/4,精度损失极小,是移动端部署的首选方案。
动态量化(快速实现,适合CPU推理)
修改代码如下:
import torch from torch.utils.mobile_optimizer import optimize_for_mobile from torchvision.models.vgg import vgg16 model = vgg16() model.load_state_dict(torch.load('./model-76-0.7754.pth')) model.eval() # 对全连接层做动态量化 quantized_model = torch.quantization.quantize_dynamic( model, {torch.nn.Linear}, dtype=torch.qint8 ) scripted_module = torch.jit.script(quantized_model) optimized_scripted_module = optimize_for_mobile(scripted_module, strip_debug_info=True) optimized_scripted_module._save_for_lite_interpreter("quantized_model.ptl")
静态量化(精度更高,需校准数据)
若对精度要求高,用静态量化,需准备少量真实样本做校准:
import torch from torch.utils.mobile_optimizer import optimize_for_mobile from torchvision.models.vgg import vgg16 model = vgg16() model.load_state_dict(torch.load('./model-76-0.7754.pth')) model.eval() # 设置量化配置(qnnpack适合移动端GPU) model.qconfig = torch.quantization.get_default_qconfig('qnnpack') torch.quantization.prepare(model, inplace=True) # 用校准数据跑一遍模型(替换为你的真实样本,数量10-100即可) calibration_data = torch.randn(20, 3, 224, 224) # 示例随机数据 with torch.no_grad(): model(calibration_data) # 完成量化转换 torch.quantization.convert(model, inplace=True) scripted_module = torch.jit.script(model) optimized_scripted_module = optimize_for_mobile(scripted_module, strip_debug_info=True) optimized_scripted_module._save_for_lite_interpreter("static_quantized_model.ptl")
2. 模型剪枝
移除模型中不重要的权重,可进一步压缩体积,对精度影响可控:
import torch from torch.nn.utils import prune from torch.utils.mobile_optimizer import optimize_for_mobile from torchvision.models.vgg import vgg16 model = vgg16() model.load_state_dict(torch.load('./model-76-0.7754.pth')) model.eval() # 对卷积层和全连接层剪枝,移除30%的权重 for name, module in model.named_modules(): if isinstance(module, torch.nn.Conv2d) or isinstance(module, torch.nn.Linear): prune.l1_unstructured(module, name='weight', amount=0.3) prune.remove(module, 'weight') # 永久移除剪枝标记 scripted_module = torch.jit.script(model) optimized_scripted_module = optimize_for_mobile(scripted_module) optimized_scripted_module._save_for_lite_interpreter("pruned_model.ptl")
3. 优化optimize_for_mobile参数
启用调试信息剥离和移动端专属优化,可减少少量体积:
optimized_scripted_module = optimize_for_mobile( scripted_module, strip_debug_info=True, # 移除调试信息 for_mobile=True # 针对移动端做优化 )
4. 替换轻量模型(可选)
如果业务允许,直接替换为MobileNet、EfficientNet-Lite等专为移动端设计的模型,体积仅为VGG16的1/10左右,同时保持不错的精度。
内容的提问来源于stack exchange,提问作者Gank
相关产品推荐
相关产品推荐

