在Seaborn子图绘制集中趋势(均值、中位数)时遇ValueError的解决
问题解决:Seaborn子图绘制直方图与统计垂直线的ValueError问题
问题背景
需要在2行3列的子图中,分别为未转换和对数转换的SalePrice、GrLivArea、GarageArea数据绘制直方图+核密度估计曲线,并为每个子图添加均值、中位数垂直线以展示集中趋势,但运行代码触发如下错误:
ValueError: The truth value of a Series is ambiguous. Use a.empty, a.bool(), a.item(), a.any() or a.all().
错误原因
- 参数类型不匹配:
df_train[list_soal].mean()和np.median(df_train[list_soal])返回的是Pandas Series(包含多列的统计值),而plt.axvline()需要单个数值作为垂直线位置,导致类型判断歧义。 - 子图定位错误:使用全局
plt.axvline()而非子图对象的ax.axvline(),无法将垂直线绘制到对应子图上。
修正后的完整代码
import pandas as pd import numpy as np import seaborn as sns %matplotlib inline import matplotlib.pyplot as plt import warnings warnings.simplefilter(action='ignore', category=FutureWarning) sns.set_theme() sns.set_style('white') df_train = pd.read_csv('https://github.com/lokalhangatt/stackoverlow/raw/refs/heads/main/train.csv') df_train = df_train.dropna(axis=1) list_soal = ['SalePrice', 'GrLivArea', 'GarageArea'] def function1(ax): ax[1].set_title('Histogram for Non-Transformed Data', fontsize=16) for i in range(len(list_soal)): col = list_soal[i] # 绘制直方图和核密度曲线 sns.histplot(df_train[col], kde=False, stat='density', bins=30, ax=ax[i]) sns.kdeplot(df_train[col], ax=ax[i]) # 计算当前列的均值、中位数 col_mean = df_train[col].mean() col_median = df_train[col].median() # 在当前子图绘制垂直线 line1 = ax[i].axvline(col_mean, color="k", linestyle="--", label="mean") line2 = ax[i].axvline(col_median, color="r", linestyle="--", label="median") # 为当前子图添加图例 ax[i].legend(handles=[line1, line2], loc=1) def function2(ax): ax[1].set_title('Histogram for Transformed Data', fontsize=16) for i in range(len(list_soal)): col = list_soal[i] # 对数转换数据 log_data = np.log10(df_train[col]) # 绘制直方图和核密度曲线 sns.histplot(log_data, kde=False, stat='density', bins=30, ax=ax[i]) sns.kdeplot(log_data, ax=ax[i]) # 计算转换后数据的均值、中位数 log_mean = log_data.mean() log_median = log_data.median() # 在当前子图绘制垂直线 line1 = ax[i].axvline(log_mean, color="k", linestyle="--", label="mean") line2 = ax[i].axvline(log_median, color="r", linestyle="--", label="median") # 为当前子图添加图例 ax[i].legend(handles=[line1, line2], loc=1) # 创建2行3列子图 fig, ax = plt.subplots(2, 3, figsize=(14,9), sharey='row') fig.subplots_adjust(hspace=0.4) function1(ax[0]) function2(ax[1]) plt.show()
关键修正点
- 循环中针对单个列计算均值、中位数,确保传入
axvline的是单个数值 - 使用子图对象的
ax[i].axvline()替代全局plt.axvline(),保证垂直线绘制在对应子图上 - 为每个子图单独添加图例,避免全局图例覆盖子图内容
- 补全了对数转换数据的均值、中位数垂直线绘制逻辑
内容的提问来源于stack exchange,提问作者lokalhangatt
相关产品推荐
相关产品推荐

