PySpark模型评估函数中MulticlassMetrics调用fMeasure()报错:缺少必填参数'label'的技术问询
解决PySpark MulticlassMetrics fMeasure()参数缺失问题
你遇到的报错fMeasure() missing 1 required positional argument: 'label'是因为PySpark的MulticlassMetrics中,fMeasure()、precision()、recall()这些方法默认是用来计算单个类别对应的指标,必须传入具体的类别标签(比如你的二分类场景中的0.0或1.0)。如果想要获取全局的整体指标(比如加权/宏观F1、精度、召回率),需要使用对应的重载方法或专门的全局指标方法。
具体解决方案
根据你的二分类场景,分两种情况处理:
1. 计算单个类别的F1、精度、召回率
直接在方法中传入目标类别标签即可,比如:
# 类别0的F1 print("F1 for class 0 = {}".format(multi_metrics.fMeasure(0.0))) # 类别1的精度 print("Precision for class 1 = {}".format(multi_metrics.precision(1.0))) # 类别0的召回率 print("Recall for class 0 = {}".format(multi_metrics.recall(0.0)))
2. 计算全局整体指标
如果你需要的是所有类别的加权/宏观综合指标,使用以下方法:
- 加权指标(按各类别样本数量加权):
weightedFMeasure()、weightedPrecision()、weightedRecall() - 宏观指标(所有类别平等加权):
macroFMeasure()、macroPrecision()、macroRecall()
修改后的完整函数示例
from pyspark.ml.evaluation import BinaryClassificationEvaluator from pyspark.mllib.evaluation import BinaryClassificationMetrics, MulticlassMetrics def print_performance_metrics(predictions): # Evaluate model with BinaryClassificationEvaluator evaluator = BinaryClassificationEvaluator(rawPredictionCol="rawPrediction") auc = evaluator.evaluate(predictions, {evaluator.metricName: "areaUnderROC"}) aupr = evaluator.evaluate(predictions, {evaluator.metricName: "areaUnderPR"}) print("auc = {}".format(auc)) print("aupr = {}".format(aupr)) # Prepare RDD for mllib metrics predictionAndLabels = predictions.select("prediction","label").rdd # Instantiate metrics objects binary_metrics = BinaryClassificationMetrics(predictionAndLabels) multi_metrics = MulticlassMetrics(predictionAndLabels) # Area under precision-recall curve (from BinaryClassificationMetrics) print("Area under PR = {}".format(binary_metrics.areaUnderPR)) # Area under ROC curve (from BinaryClassificationMetrics) print("Area under ROC = {}".format(binary_metrics.areaUnderROC)) # Accuracy print("Accuracy = {}".format(multi_metrics.accuracy)) # Confusion Matrix print("Confusion Matrix:\n{}".format(multi_metrics.confusionMatrix())) ### 修复后的F1、Precision、Recall指标 ### # 全局加权F1 print("Weighted F1 = {}".format(multi_metrics.weightedFMeasure())) # 全局宏观F1 print("Macro F1 = {}".format(multi_metrics.macroFMeasure())) # 类别0的F1 print("F1 for class 0 = {}".format(multi_metrics.fMeasure(0.0))) # 全局加权精度 print("Weighted Precision = {}".format(multi_metrics.weightedPrecision())) # 类别1的精度 print("Precision for class 1 = {}".format(multi_metrics.precision(1.0))) # 全局加权召回率 print("Weighted Recall = {}".format(multi_metrics.weightedRecall())) # 类别0的召回率 print("Recall for class 0 = {}".format(multi_metrics.recall(0.0))) # FPR for class 0 print("FPR for class 0 = {}".format(multi_metrics.falsePositiveRate(0.0))) # TPR for class 0 print("TPR for class 0 = {}".format(multi_metrics.truePositiveRate(0.0)))
额外说明
在你的二分类场景中,其实BinaryClassificationMetrics也能提供部分二分类专属指标,但MulticlassMetrics的优势是可以同时查看单个类别的细节指标,适合需要深入分析类别表现的场景。
内容的提问来源于stack exchange,提问作者Adem Arslan
相关产品推荐
相关产品推荐

