加载Pickle保存的ML模型报错:Buffer dtype mismatch, expected 'ITYPE_t' but got 'long long'
ValueError: Buffer dtype mismatch, expected 'ITYPE_t' but got 'long long'
报错根本原因
该错误属于跨平台模型序列化的兼容性问题,和你使用的序列化工具(pickle/dill/joblib/ubjson)没有关联,问题根源在KNN模型的底层实现逻辑:
- scikit-learn的KNN分类器底层基于Cython开发,
ITYPE_t是Cython层定义的平台相关整数类型,32位/64位环境、不同版本的scikit-learn对该类型的底层映射规则存在差异 - 你用来训练保存模型的设备和加载模型的设备存在至少一项以下差异:
- Python架构不同(一台为32位Python、另一台为64位Python)
- 操作系统不同(比如Windows环境训练、Linux/macOS环境加载,反之亦然)
- scikit-learn、numpy等依赖库的版本差异过大,底层整数类型定义有变动
可行解决方案
- 方案1:对齐两端环境
确保训练端和加载端的运行环境完全匹配:- scikit-learn、numpy、scipy的版本号完全一致
- Python的架构(32位/64位)完全一致
- 尽量使用同类型操作系统,避免跨操作系统序列化模型
- 方案2:手动导出可移植参数,避免序列化完整模型
不直接序列化整个KNN模型对象,手动提取训练后的核心参数,在目标设备上重新初始化KNN实例加载参数,参考代码如下:
训练端保存参数:
加载端重建模型:import numpy as np # 提取KNN训练后的核心参数和配置 knn_params = { "n_neighbors": src_model.n_neighbors, "metric": src_model.metric, "weights": src_model.weights, "_fit_X": src_model._fit_X, "_y": src_model._y, "classes_": src_model.classes_ } # 保存参数到可移植格式文件 np.savez("knn_params.npz", **knn_params)from sklearn.neighbors import KNeighborsClassifier import numpy as np params = np.load("knn_params.npz", allow_pickle=True) # 初始化KNN实例 loaded_model = KNeighborsClassifier( n_neighbors=int(params["n_neighbors"]), metric=str(params["metric"]), weights=str(params["weights"]) ) # 手动填充训练得到的属性 loaded_model._fit_X = params["_fit_X"] loaded_model._y = params["_y"] loaded_model.classes_ = params["classes_"] - 方案3:使用跨平台兼容的ONNX格式导出
如需频繁跨设备部署,可导出为ONNX通用格式,完全脱离平台和sklearn版本限制,参考代码如下:
训练端导出ONNX模型:
加载端用ONNX Runtime推理:from skl2onnx import convert_sklearn from skl2onnx.common.data_types import FloatTensorType # 按实际输入特征维度修改n的取值 initial_type = [("float_input", FloatTensorType([None, n]))] onnx_model = convert_sklearn(src_model, initial_types=initial_type) with open("knn.onnx", "wb") as f: f.write(onnx_model.SerializeToString())import onnxruntime as rt import numpy as np sess = rt.InferenceSession("knn.onnx") input_name = sess.get_inputs()[0].name label_name = sess.get_outputs()[0].name # 执行推理,X_test为待预测的特征数组 pred = sess.run([label_name], {input_name: X_test.astype(np.float32)})
内容的提问来源于stack exchange,提问作者Mohamed Ahmed
相关产品推荐
相关产品推荐

