如何高效保存PyTorch训练的协同过滤模型以缩减文件体积?
高效压缩PyTorch/fastAI协同过滤模型文件的实用方案
1. 只存模型权重,别存整个对象
PyTorch默认保存完整模型实例(含结构、权重、关联组件),但协同过滤模型的核心是Embedding等层的权重,仅保存权重字典能砍掉大部分冗余体积:
原生PyTorch操作
# 保存权重 torch.save(model.state_dict(), 'model_weights.pth') # 加载:先重建模型结构,再载入权重 model = YourCollabFilterModel() model.load_state_dict(torch.load('model_weights.pth')) model.eval()
fastAI Learner适配
你的export方法已经清空了dataloader和优化器,还可以更进一步只导出权重:
# 导出权重字典 torch.save(learn.model.state_dict(), 'model_weights.pth') # 加载时先恢复Learner结构,再重载权重 learn = load_learner('export.pkl') learn.model.load_state_dict(torch.load('model_weights.pth'))
2. 启用PyTorch内置压缩序列化
torch.save自带_use_new_zipfile_serialization=True参数,专门针对张量数据做了压缩优化,比普通pickle或第三方工具更适配PyTorch,压缩效果明显:
原生PyTorch用例
torch.save(model.state_dict(), 'model_compressed.pth', _use_new_zipfile_serialization=True)
改造你的fastAI export方法
直接在原有代码里加这个参数即可:
def export(self:Learner, fname='export.pkl', pickle_module=pickle, pickle_protocol=2): "Export the content of `self` without the items and the optimizer state for inference" if rank_distrib(): return # don't export if child proc self._end_cleanup() old_dbunch = self.dls self.dls = self.dls.new_empty() state = self.opt.state_dict() if self.opt is not None else None self.opt = None with warnings.catch_warnings(): warnings.simplefilter("ignore") # 启用PyTorch原生压缩 torch.save(self, self.path/fname, pickle_module=pickle_module, pickle_protocol=pickle_protocol, _use_new_zipfile_serialization=True) self.create_opt() if state is not None: self.opt.load_state_dict(state) self.dls = old_dbunch
3. 换用高效压缩算法
compress-pickle增大体积大概率是用了不合适的默认算法(比如gzip对二进制张量压缩率低),换成lz4或zstd这类适合二进制数据的算法,压缩速度和率都更优:
import lz4.frame # 保存权重并压缩 weights = model.state_dict() with lz4.frame.open('model_weights.lz4', 'wb') as f: torch.save(weights, f) # 加载 with lz4.frame.open('model_weights.lz4', 'rb') as f: weights = torch.load(f) model.load_state_dict(weights)
4. 模型量化(精度允许的情况下)
如果你的场景对模型精度要求不苛刻,用PyTorch的动态量化把模型权重从32位浮点转成8位整数,体积直接砍到原来的1/4,还能加快推理:
# 针对Embedding和Linear层做动态量化(协同过滤模型的核心层) quantized_model = torch.quantization.quantize_dynamic( model, {torch.nn.Embedding, torch.nn.Linear}, dtype=torch.qint8 ) # 保存量化后的模型,配合内置压缩 torch.save(quantized_model.state_dict(), 'quantized_model.pth', _use_new_zipfile_serialization=True)
内容的提问来源于stack exchange,提问作者Silver Light
相关产品推荐
相关产品推荐

