使用MLFlow保存PyTorch模型时出现Can't get attribute Net错误如何解决
报错原因
- 核心为PyTorch序列化(pickle)的作用域特性导致:你在Jupyter Notebook中定义的
Net类属于当前Notebook的__main__作用域,MLFlow执行依赖推断逻辑时会启动独立子进程加载已序列化的模型,子进程的__main__模块是MLFlow自带的_capture_modules.py,不存在你自定义的Net类,因此反序列化失败。 - 报错开头的CUDA版本警告不影响本次错误,仅为MLFlow默认将带
+cu111的本地版本替换为PyPI可安装的公开版本号,可后续按需调整。
解决方案
以下方案任选一种即可:
方案1:关闭MLFlow自动依赖推断
在log_model时手动指定依赖,跳过自动推断步骤,不会触发子进程加载模型的逻辑:
mlflow.pytorch.log_model( net, artifact_path="model", pickle_module=pickle, # 也可以传入自定义requirements.txt的路径 pip_requirements=["torch==1.9.0+cu111", "torchvision==0.10.0+cu111"] )
方案2:将Net类抽为独立Python文件
不要在Notebook内定义Net类,单独写入例如model.py文件,在Notebook中通过from model import Net导入后再实例化训练。此时序列化的模型关联的是model.Net而非__main__.Net,子进程加载时只要sys路径包含model.py所在目录即可找到类定义。
方案3:仅保存模型权重(state_dict)
不直接保存整个模型对象,仅保存权重,同时将模型结构代码随MLFlow制品一同上传,后续加载时先实例化Net类再加载权重:
# 保存阶段 mlflow.log_dict(net.state_dict(), "model_weights.pth") # 加载阶段 net = Net() net.load_state_dict(mlflow.load_dict("runs:/<RUN_ID>/model_weights.pth"))
方案4:临时修改Net类的模块属性
保存模型前手动指定Net类的模块名,避免关联到__main__:
Net.__module__ = "model" net = Net() # 后续正常调用log_model即可
内容的提问来源于stack exchange,提问作者ucsky
相关产品推荐
相关产品推荐

