You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)

错误原因分析

  1. 方法2完全不适用:CatBoostRegressionModel.load()是用来加载Spark环境下训练并通过save()方法保存的模型,会生成包含metadata文件的目录结构。而Python中save_model()保存的是原生CBM单文件,两者格式不兼容,因此该方法无法使用。

  2. 方法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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 01:32:05