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

Spark GBTClassifier训练后预测准确率100%异常问题求助

Why Your GBTClassifier Shows "100% Accuracy" (And How to Fix It)

Hey there, let's dig into why your SparkML GBTClassifier is returning a misleading "perfect" accuracy score, and walk through concrete fixes to build a actually useful model.

Root Causes

First, let's break down what's happening here:

  1. Extreme Class Imbalance: Your training set has a massive skew—class 0 makes up ~97% of the data, while class 1 is only ~3%. GBT (like most tree-based models) optimizes for minimizing overall loss, so it's way easier for the model to just predict every sample as class 0. This gives an extremely high accuracy score, but it's completely useless because the model never learns to identify class 1.
  2. Misleading Evaluation Metric: Accuracy is a terrible choice for imbalanced datasets. It only measures overall correct predictions, which is dominated by the majority class. A model that does nothing but guess class 0 will already have ~97% accuracy here—so hitting "100%" just means your test set follows the same skew, and the model didn't miss any class 0 samples (or the test set has no class 1 samples at all).

Fixes to Try

Let's tackle this from three angles: balancing your data, adjusting the model, and using meaningful metrics.

1. Balance Your Training Data

You can adjust the class distribution to force the model to pay attention to the minority class:

  • Over-Sample the Minority Class: Use RandomOverSampler to duplicate class 1 samples until the distribution is balanced. This preserves all majority class data.
    from pyspark.ml.feature import RandomOverSampler
    
    ros = RandomOverSampler(labelCol=labelCol, seed=420)
    Xtrain_balanced = ros.fit(Xtrain).transform(Xtrain)
    # Verify the new distribution
    Xtrain_balanced.select(labelCol).groupBy(labelCol).count().show()
    
  • Under-Sample the Majority Class: Use RandomUnderSampler to randomly remove class 0 samples. This reduces dataset size but avoids overfitting to duplicated minority samples.
  • Hybrid Sampling: Combine both (e.g., under-sample class 0 to 2x the size of class 1, then over-sample class 1 to match) to balance data while minimizing information loss.

2. Assign Class Weights to the Model

Instead of modifying the data, you can tell the GBTClassifier to weight minority class errors more heavily using the weightCol parameter:
First, calculate inverse-frequency weights for each class (so minority samples have higher weight):

from pyspark.sql.functions import when, lit

total_samples = 631608 + 18428
class_0_count = 631608
class_1_count = 18428

# Add a weight column: weight = total_samples / (2 * class_sample_count)
Xtrain = Xtrain.withColumn(
    "weight",
    when(Xtrain[labelCol] == 1, lit(total_samples / (2 * class_1_count)))
    .otherwise(lit(total_samples / (2 * class_0_count)))
)

Then update your GBTClassifier to use this weight column:

tr = GBTClassifier(
    labelCol=labelCol,
    featuresCol="features",
    maxIter=30,
    maxDepth=5,
    seed=420,
    weightCol="weight"  # Add this line
)

This makes the model penalize misclassifying class 1 samples much more, forcing it to learn patterns that identify the minority class.

3. Use Meaningful Evaluation Metrics

Stop relying on accuracy—switch to metrics that reflect how well the model performs on the minority class:

  • AUC-PR (Area Under Precision-Recall Curve): Far more sensitive to imbalanced data than AUC-ROC, as it focuses on the minority class's performance.
  • Precision, Recall, F1-Score: Precision measures how many predicted class 1 samples are actually class 1; recall measures how many real class 1 samples the model catches; F1 is the harmonic mean of the two.

Here's how to calculate these in Spark:

from pyspark.ml.evaluation import BinaryClassificationEvaluator, MulticlassClassificationEvaluator

predictions = model.transform(Xtest)

# AUC-PR (best for imbalanced data)
pr_evaluator = BinaryClassificationEvaluator(labelCol=labelCol, metricName="areaUnderPR")
print(f"AUC-PR Score: {pr_evaluator.evaluate(predictions):.4f}")

# Precision, Recall, F1 for class 1
multi_evaluator = MulticlassClassificationEvaluator(
    labelCol=labelCol,
    predictionCol="prediction"
)
precision = multi_evaluator.evaluate(predictions, {multi_evaluator.metricName: "precisionByLabel", multi_evaluator.metricLabel: 1})
recall = multi_evaluator.evaluate(predictions, {multi_evaluator.metricName: "recallByLabel", multi_evaluator.metricLabel: 1})
f1 = multi_evaluator.evaluate(predictions, {multi_evaluator.metricName: "f1ByLabel", multi_evaluator.metricLabel: 1})

print(f"Class 1 Precision: {precision:.4f}")
print(f"Class 1 Recall: {recall:.4f}")
print(f"Class 1 F1 Score: {f1:.4f}")

# Check confusion matrix to see actual prediction breakdown
from pyspark.sql.functions import count
confusion_matrix = predictions.groupBy(labelCol, "prediction").agg(count("*").alias("count")).orderBy(labelCol, "prediction")
confusion_matrix.show()

Final Notes

Start with class weights first—it's simpler than resampling and preserves all your data. Then use AUC-PR and class-specific metrics to evaluate your model, not accuracy. You'll likely see the "100% accuracy" disappear, but you'll get a model that actually does what you need: identify class 1 samples.

内容的提问来源于stack exchange,提问作者Dimon Buzz

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 10:02:38