PySpark 2.4.4中GBT交叉验证报错:'GBTClassificationModel'无fitMultiple属性
解决PySpark 2.4.4中GBT交叉验证的
fitMultiple属性错误 我来帮你排查这个问题!你遇到的AttributeError: 'GBTClassificationModel' object has no attribute 'fitMultiple'报错,核心原因是传入CrossValidator的gbt是已经训练好的模型实例,而不是未训练的估计算法类。
问题根源
在PySpark的MLlib体系里:
GBTClassifier是Estimator(估计算法):它是用来训练模型的组件,自带fit()和fitMultiple()方法(后者是CrossValidator并行训练参数组合所需的接口)。GBTClassificationModel是Model(训练后的模型):它是Estimator训练后的产物,只有transform()等用于预测的方法,没有fitMultiple。
你之前单独运行GBT无异常,应该是先通过gbt.fit(train)得到了GBTClassificationModel,之后误把这个训练好的模型当成Estimator传给了CrossValidator,才触发了这个错误。
修正方案
确保你的gbt是GBTClassifier的实例(未训练的Estimator),而不是训练后的模型。修正后的代码示例:
from pyspark.ml.classification import GBTClassifier from pyspark.ml.tuning import ParamGridBuilder, CrossValidator from pyspark.ml.evaluation import BinaryClassificationEvaluator # 第一步:创建未训练的GBTClassifier实例(Estimator) gbt = GBTClassifier(labelCol="label", featuresCol="features") # 根据你的数据调整列名 # 构建参数网格 paramGrid = ParamGridBuilder()\ .addGrid(gbt.maxDepth, [2, 4, 6])\ .addGrid(gbt.maxBins, [20, 60])\ .addGrid(gbt.maxIter, [10, 20])\ .build() evaluator = BinaryClassificationEvaluator() cv = CrossValidator(estimator=gbt, estimatorParamMaps=paramGrid, evaluator=evaluator) # 现在执行交叉验证就不会报错了 cvModel = cv.fit(train) predictions = cvModel.transform(test) evaluator.evaluate(predictions)
额外说明
PySpark 2.4.4中的GBTClassifier是支持fitMultiple方法的,这是CrossValidator并行训练多个参数组合的基础。只要你传入的是正确的Estimator类型,就可以正常运行交叉验证。
内容的提问来源于stack exchange,提问作者griffinleow
相关产品推荐
相关产品推荐

