Spark GBTClassifier训练后预测准确率100%异常问题求助
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:
- 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.
- 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
RandomOverSamplerto 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
RandomUnderSamplerto 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

