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

如何基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 09:22:38