如何基于PySpark DataFrame绘制R²值并解决可视化报错
解决PySpark DataFrame传入sklearn r2_score的报错及散点图标注R²
报错原因
sklearn的r2_score函数仅支持类数组输入(如numpy数组、pandas Series、Python列表),PySpark DataFrame是分布式数据结构,无法直接作为参数传入,这才触发InvalidParameterError提示"y_true需为类数组对象"。
解决方案步骤
1. 正确计算R²:将PySpark数据转为本地类数组结构
先把PySpark DataFrame中需要的真实值、预测值列拉取到本地,转为pandas或numpy格式,再传入r2_score:
from sklearn.metrics import r2_score # 假设你的SparkRatingEvaluation类输出的PySpark DataFrame为spark_df,包含true_rating和predicted_rating列 # 方式1:转为pandas DataFrame后提取列 pd_df = spark_df.select("true_rating", "predicted_rating").toPandas() y_true = pd_df["true_rating"] y_pred = pd_df["predicted_rating"] # 方式2:直接转为Python列表 y_true = spark_df.select("true_rating").rdd.flatMap(lambda x: x).collect() y_pred = spark_df.select("predicted_rating").rdd.flatMap(lambda x: x).collect() # 计算R² r2 = r2_score(y_true, y_pred)
注意:如果PySpark数据量极大,直接
toPandas()或collect()可能导致内存溢出,建议先采样:
采样10%的数据(可根据实际调整比例)
sampled_spark_df = spark_df.sample(withReplacement=False, fraction=0.1, seed=42)
pd_df = sampled_spark_df.select("true_rating", "predicted_rating").toPandas()
### 2. 绘制散点图并标注R² 用matplotlib绘制真实值vs预测值的散点图,并用`annotate`方法标注计算好的R²值: ```python import matplotlib.pyplot as plt plt.figure(figsize=(8, 6)) # 绘制散点图 plt.scatter(pd_df["true_rating"], pd_df["predicted_rating"], alpha=0.5, s=20) # 添加理想预测对角线 plt.plot([pd_df["true_rating"].min(), pd_df["true_rating"].max()], [pd_df["true_rating"].min(), pd_df["true_rating"].max()], 'r--', label="Ideal Prediction") # 标注R²值 plt.annotate(f"$R^2 = {r2:.4f}$", xy=(0.05, 0.95), xycoords='axes fraction', fontsize=12, bbox=dict(facecolor='white', alpha=0.8)) # 设置图表标签 plt.xlabel("True Rating") plt.ylabel("Predicted Rating") plt.title("True vs Predicted Ratings") plt.legend() plt.show()
完整代码示例
from sklearn.metrics import r2_score import matplotlib.pyplot as plt # 假设已通过SparkRatingEvaluation类得到结果PySpark DataFrame spark_df # 1. 处理数据并计算R² sampled_spark_df = spark_df.sample(withReplacement=False, fraction=0.1, seed=42) pd_df = sampled_spark_df.select("true_rating", "predicted_rating").toPandas() r2 = r2_score(pd_df["true_rating"], pd_df["predicted_rating"]) # 2. 绘制散点图并标注 plt.figure(figsize=(8, 6)) plt.scatter(pd_df["true_rating"], pd_df["predicted_rating"], alpha=0.5, s=20) plt.plot([pd_df["true_rating"].min(), pd_df["true_rating"].max()], [pd_df["true_rating"].min(), pd_df["true_rating"].max()], 'r--', label="Ideal Prediction") plt.annotate(f"$R^2 = {r2:.4f}$", xy=(0.05, 0.95), xycoords='axes fraction', fontsize=12, bbox=dict(facecolor='white', alpha=0.8)) plt.xlabel("True Rating") plt.ylabel("Predicted Rating") plt.title("True vs Predicted Ratings") plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者tran su
相关产品推荐
相关产品推荐

