Python训练含分类特征的CatBoost模型如何在Spark中加载?
解决Python训练的CatBoost CBM模型在Spark中加载失败的问题
问题背景
Python中用CatBoost Regressor训练含分类特征的回归模型,因ONNX不支持CatBoost分类特征、PMML需独热编码导致模型膨胀,选择保存为原生CBM格式。将模型上传至HDFS后,Spark中两种加载方式均报错:
方法1(loadNativeModel)报错
代码:
import ai.catboost.spark._ val loadedModel = CatBoostRegressionModel.loadNativeModel("/path/to/model.cbm")
错误信息:
ai.catboost.CatBoostError: /src/catboost/catboost/libs/model/model_import_interface.h:19: Model file doesn't exist: /path/to/model.cbm at ru.yandex.catboost.spark.catboost4j_spark.core.src.native_impl.native_implJNI.ReadModel__SWIG_0(Native Method) at ru.yandex.catboost.spark.catboost4j_spark.core.src.native_impl.native_impl.ReadModel(native_impl.java:193) at ai.catboost.spark.CatBoostRegressionModel$.loadNativeModel(CatBoostRegressor.scala:145)
方法2(load)报错
代码:
import ai.catboost.spark._ val loadedModel = CatBoostRegressionModel.load("/path/to/parent_directory")
错误信息:
org.apache.hadoop.mapred.InvalidInputException: Input path does not exist: hdfs://path/to/parent_directory/metadata at org.apache.hadoop.mapred.LocatedFileStatusFetcher.getFileStatuses(LocatedFileStatusFetcher.java:156) at org.apache.hadoop.mapred.FileInputFormat.listStatus(FileInputFormat.java:247) at org.apache.hadoop.mapred.FileInputFormat.getSplits(FileInputFormat.java:325)
错误原因分析
方法2完全不适用:
CatBoostRegressionModel.load()是用来加载Spark环境下训练并通过save()方法保存的模型,会生成包含metadata文件的目录结构。而Python中save_model()保存的是原生CBM单文件,两者格式不兼容,因此该方法无法使用。方法1的问题:
loadNativeModel()默认读取本地文件系统的路径,而你的模型存储在HDFS上,直接传入HDFS路径会导致找不到文件。
解决方案
方案1:将HDFS模型下载到集群本地节点加载
先通过HDFS命令将模型文件下载到Spark集群每个节点的相同本地路径(确保所有节点都能访问到):
hdfs dfs -get hdfs://path/to/model.cbm /local/node/path/model.cbm
然后在Spark中用本地路径加载:
import ai.catboost.spark._ val loadedModel = CatBoostRegressionModel.loadNativeModel("/local/node/path/model.cbm")
方案2:通过Hadoop API读取HDFS模型字节流加载
无需下载文件,直接通过Hadoop FileSystem API读取HDFS上的模型文件为字节数组,再调用loadNativeModel的重载方法加载:
import ai.catboost.spark._ import org.apache.hadoop.fs.{FileSystem, Path} import org.apache.spark.SparkContext // 获取Spark上下文和HDFS文件系统实例 val sc = SparkContext.getOrCreate() val hdfsConf = sc.hadoopConfiguration val fs = FileSystem.get(hdfsConf) // 读取HDFS上的模型文件为字节数组 val modelHdfsPath = new Path("hdfs://path/to/model.cbm") val inputStream = fs.open(modelHdfsPath) val modelBytes = Stream.continually(inputStream.read()).takeWhile(_ != -1).map(_.toByte).toArray inputStream.close() // 加载模型 val loadedModel = CatBoostRegressionModel.loadNativeModel(modelBytes)
额外注意事项
- 版本一致性:确保Python中使用的CatBoost版本与Spark依赖的
catboost4j-spark版本完全一致,版本不匹配会导致模型加载失败或运行异常。
内容的提问来源于stack exchange,提问作者return007
相关产品推荐
相关产品推荐

