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

Spark ML 2.0多分类随机森林模型AUC计算方法问询

Solution for Multi-Class AUC in Spark 2.0

Since Spark 2.0's MulticlassClassificationEvaluator doesn’t support AUC directly for multi-class tasks, we can use a One-vs-Rest (OvR) strategy to compute AUC for each class individually, then aggregate these values (like a weighted average based on class frequency) to get an overall multi-class AUC score. Here’s how to adjust your existing code:

Step-by-Step Implementation

from pyspark.ml.classification import RandomForestClassifier
from pyspark.ml.evaluation import MulticlassClassificationEvaluator, BinaryClassificationEvaluator
from pyspark.sql.functions import col, when

# Your existing training workflow stays intact
max_depth = model_params['max_depth']
num_trees = model_params['num_trees']

rf = RandomForestClassifier(labelCol="label", featuresCol="features", impurity="gini", 
                            featureSubsetStrategy="all", numTrees=num_trees, maxDepth=max_depth)
model_fit = rf.fit(training_data)
transformed = model_fit.transform(test_data)

# Calculate accuracy (your original logic)
if model_params['calc_matrix'] is True:
    evaluator_acc = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction", metricName="accuracy")
    accuracy = evaluator_acc.evaluate(transformed)
    print(f"RF Overall Accuracy = {accuracy}, numTrees = {num_trees}, maxDepth = {max_depth}")

# Calculate multi-class AUC using One-vs-Rest approach
# 1. Get all unique classes from your label column
unique_classes = transformed.select("label").distinct().rdd.flatMap(lambda x: x).collect()
unique_classes.sort()

# 2. Initialize binary evaluator for AUC calculations
evaluator_auc = BinaryClassificationEvaluator(rawPredictionCol="rawPrediction", 
                                               labelCol="binary_label", 
                                               metricName="areaUnderROC")

total_weighted_auc = 0.0
total_samples = transformed.count()

for cls in unique_classes:
    # Create a binary label: current class = 1 (positive), all others = 0 (negative)
    binary_labeled_df = transformed.withColumn("binary_label", when(col("label") == cls, 1).otherwise(0))
    
    # Compute AUC for this class vs all others
    class_auc = evaluator_auc.evaluate(binary_labeled_df)
    
    # Calculate weight based on class frequency (adjust to equal weights if needed)
    class_sample_count = binary_labeled_df.filter(col("binary_label") == 1).count()
    class_weight = class_sample_count / total_samples
    
    total_weighted_auc += class_auc * class_weight
    print(f"AUC for class {cls}: {class_auc}, weight: {class_weight:.4f}")

print(f"Weighted Average Multi-Class AUC = {total_weighted_auc:.4f}")

Key Details:

  • One-vs-Rest Logic: For each class, we frame the problem as a binary classification task (current class vs all others) and compute AUC for each scenario.
  • Weighted Aggregation: Using class frequencies as weights ensures classes with more samples have a proportional impact on the overall AUC. If you prefer equal importance for all classes, replace class_weight with 1/len(unique_classes).
  • RawPredictionCol: Spark’s BinaryClassificationEvaluator uses raw model scores (rawPrediction) instead of probabilities for consistent AUC calculation.

内容的提问来源于stack exchange,提问作者E B

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:01:06