You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效保存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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 04:52:05