求助:使用Scikit-learn模型创建AI Platform版本时加载失败
我之前也碰到过一模一样的情况,结合你的描述(模型大小合规但加载报错),可以从这几个方向排查和解决:
1. 优先检查scikit-learn版本兼容性
这是最常见的触发原因——AI Platform预测环境使用的scikit-learn版本,和你本地训练模型时的版本不匹配。比如你用0.24.2训练模型,但AI Platform默认runtime用的是更旧的0.22.2,就会导致模型加载失败。
- 解决步骤:
- 先查本地训练时的版本:
pip show scikit-learn - 创建模型版本时指定对应兼容的runtime版本,比如你用的是scikit-learn 0.24.2,就搭配
runtime-version 2.5(不同runtime对应固定的框架版本,只要保证两者匹配即可) - 命令示例:
gcloud ai-platform versions create YOUR_VERSION_NAME \ --model YOUR_MODEL_NAME \ --origin gs://YOUR_BUCKET_PATH/TO/MODEL_FOLDER/ \ --runtime-version 2.5 \ --framework scikit-learn \ --python-version 3.7
- 先查本地训练时的版本:
2. 确认模型存储路径结构是否正确
AI Platform要求模型文件必须直接放在你指定的GCS origin路径下,不能嵌套子文件夹。比如你的模型存在gs://my-bucket/rf-model/model.joblib,那么创建版本时的--origin必须是gs://my-bucket/rf-model/,而不是指向具体的model.joblib文件。
从报错路径/tmp/model/a0001/model.joblib来看,AI Platform在临时目录的子文件夹里找模型,大概率是你的origin路径设置有误,导致它没正确提取到模型文件。
3. 验证模型本身是否完好
先把GCS里的model.joblib下载到本地,用和训练时一模一样的scikit-learn版本尝试加载并测试:
from joblib import load import numpy as np # 加载模型 model = load('model.joblib') # 用样本数据测试预测 test_sample = np.array([[1.2, 3.4, 5.6, 7.8]]) model.predict(test_sample)
如果本地加载失败,说明模型在保存过程中损坏了,需要重新训练并保存;如果本地完全正常,那问题肯定出在AI Platform的环境配置上。
4. 检查模型保存方式是否合规
确保你用的是scikit-learn原生的joblib保存方式,不要混合使用pickle,也不要在模型中包含自定义未序列化的组件(比如自定义预处理函数):
from sklearn.ensemble import RandomForestClassifier from joblib import dump # 训练模型 rf_model = RandomForestClassifier(n_estimators=100) rf_model.fit(X_train, y_train) # 正确保存方式 dump(rf_model, 'model.joblib')
如果模型里有自定义组件,需要把这些组件的代码打包成可导入的模块,确保AI Platform环境能正确加载。
内容的提问来源于stack exchange,提问作者Ajay Deshpande

