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_weightwith1/len(unique_classes). - RawPredictionCol: Spark’s
BinaryClassificationEvaluatoruses raw model scores (rawPrediction) instead of probabilities for consistent AUC calculation.
内容的提问来源于stack exchange,提问作者E B
相关产品推荐
相关产品推荐

