Python PyTorch训练的神经网络无法在R torch包中加载的问题排查
解决Python PyTorch模型无法在R torch包中加载的问题
问题根源
你的Python PyTorch版本(2.2.1+cu121)远高于R torch包依赖的libtorch版本(2.0.1),PyTorch的序列化格式在跨大版本间存在兼容性差异,这是两次报错的核心原因:
- 首次报错
rawToChar(raw_json) : embedded nul in string:新版本PyTorch默认的序列化格式包含旧版本R torch无法解析的内容 - 二次报错
Unpickler found unknown type torch.nn.modules.container.Sequential:版本差异导致旧libtorch无法识别新版本序列化的模型类
解决方案
方案1:对齐PyTorch版本(最稳妥)
- 降级Python端PyTorch:将Python的PyTorch版本降级到2.0.x系列,与R torch依赖的libtorch版本匹配
pip install torch==2.0.1 torchvision==0.15.2 torchaudio==2.0.2 --index-url https://download.pytorch.org/whl/cu118 - 重新保存模型:使用传统pickle序列化方式(不要加
_use_new_zipfile_serialization=True)torch.save(model.state_dict(), "test_save_pytorchNN.pth") - R端加载:先定义与Python完全一致的模型结构,再加载状态字典
library(torch) # 示例:定义和Python中相同的模型结构 model <- nn_sequential( nn_linear(784, 256), nn_relu(), nn_linear(256, 10) ) # 加载状态字典 model$load_state_dict(torch_load("test_save_pytorchNN.pth"))
方案2:升级R torch包
- 升级R torch包及依赖的libtorch:
install.packages("torch") torch::install_torch(version = "2.2.1") # 指定与Python匹配的版本 - R端加载:同样先定义与Python一致的模型结构,再加载状态字典(支持Python端的新旧序列化格式)
关键注意事项
- 必须保证R中定义的模型结构与Python中的完全一致:包括层的类型、顺序、输入输出维度、激活函数等,否则加载状态字典会失败
- 避免跨大版本序列化PyTorch模型,尽量保持两端PyTorch主版本号一致
内容的提问来源于stack exchange,提问作者Bastien
相关产品推荐
相关产品推荐

