PySpark中Logistic Regression模型评估指标代码编写与报错排查
PySpark Logistic Regression 性能评估代码修复方案
问题根源
你的代码报错核心原因有两个:
- 原始数据缺少分类模型必需的标签列(Logistic Regression是监督学习,必须依赖标签完成训练与评估)
- 未提前定义
evaluator_accuracy、evaluator_precision等评估器对象,直接调用会抛出未定义错误
完整可运行修复代码
以下是从数据预处理、模型训练到指标评估的全流程代码:
1. 导入依赖模块
from pyspark.sql import SparkSession from pyspark.ml.feature import StringIndexer, VectorAssembler from pyspark.ml.classification import LogisticRegression from pyspark.ml.evaluation import MulticlassClassificationEvaluator
2. 初始化Session与数据预处理
先给原始数据添加二分类标签(这里以薪资是否超过50000作为划分标准,1表示超过,0表示未超过):
spark = SparkSession.builder.appName("LREvaluation").getOrCreate() # 原始数据 data = spark.createDataFrame([ (0, 18.0, "male", 5.0, 35000), (1, 20.0, "female", 3.0, 45000), (2, 22.0, "male", 8.0, 58000), (3, 25.0, "female", 2.0, 62000), ], ["id", "age", "gender", "experience", "salary"]) # 添加标签列 data_with_label = data.withColumn("label", (data["salary"] > 50000).cast("integer"))
3. 特征工程处理
将分类特征gender转为数值型,再组装成模型需要的特征向量:
# 处理分类特征gender indexer = StringIndexer(inputCol="gender", outputCol="gender_indexed") data_indexed = indexer.fit(data_with_label).transform(data_with_label) # 组装特征向量 assembler = VectorAssembler( inputCols=["age", "gender_indexed", "experience", "salary"], outputCol="features" ) final_data = assembler.transform(data_indexed).select("label", "features")
4. 划分训练/测试集
train_data, test_data = final_data.randomSplit([0.7, 0.3], seed=42)
5. 训练LR模型并生成预测结果
lr = LogisticRegression(featuresCol="features", labelCol="label") lr_model = lr.fit(train_data) # 生成测试集预测结果 predictions = lr_model.transform(test_data)
6. 定义评估器并计算指标
无需创建多个评估器,使用MulticlassClassificationEvaluator即可一次性计算所有需要的指标:
# 初始化多分类评估器 evaluator = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction") # 计算各指标 accuracy = evaluator.evaluate(predictions, {evaluator.metricName: "accuracy"}) precision = evaluator.evaluate(predictions, {evaluator.metricName: "weightedPrecision"}) recall = evaluator.evaluate(predictions, {evaluator.metricName: "weightedRecall"}) f1 = evaluator.evaluate(predictions, {evaluator.metricName: "f1"}) # 打印结果 print(f"Accuracy: {accuracy:.2f}") print(f"Precision: {precision:.2f}") print(f"Recall: {recall:.2f}") print(f"F1 Score: {f1:.2f}")
关键注意事项
- 若为二分类场景,
BinaryClassificationEvaluator主要支持AUC、ROC等指标,精度/召回/F1用MulticlassClassificationEvaluator更直接 - 确保
predictions数据框包含label(真实标签)和prediction(模型预测值)两列,这是评估器的必需输入列 - 你的样本量极小(仅4条),划分训练测试集后可能出现某类别样本缺失,导致指标计算报错,建议增加样本量或调整划分比例
内容的提问来源于stack exchange,提问作者Sabrina Cheruiyot
相关产品推荐
相关产品推荐

