如何在子图中绘制多个带相关系数标注的Seaborn Jointplot
问题原因
sns.jointplot() 设计为独立生成包含「主关联图+X轴边缘分布图+Y轴边缘分布图」的专属画布,不支持通过ax参数嵌入到plt.subplots()预创建的子图网格中,你传入的ax参数不会被接口识别,因此每次循环调用都会生成单独的图表,无法合并到预设布局里。
解决方案
根据你需要的最终效果可选择两种实现方式:
方案1:网格共享边缘分布(推荐,更简洁)
你的需求是固定Y轴变量列表为rt、X轴变量列表为nvars的网格布局,直接用seaborn.PairGrid即可原生实现,不需要手动创建子图,每列共享X轴边缘、每行共享Y轴边缘:
import numpy as np import pandas as pd import scipy.stats as stats import seaborn as sns import matplotlib.pyplot as plt ncols=['ra','rb','a','b','c','d'] df=pd.DataFrame(np.random.rand(100,len(ncols)),columns=ncols) nvars=['a','b','c','d'] rt=['a','b'] # 创建2行4列的配对网格,对应rt的长度为行,nvars的长度为列 g = sns.PairGrid(df, x_vars=nvars, y_vars=rt, height=3) # 主区域画回归图,对应原jointplot的kind='reg' g.map(sns.regplot, line_kws={"color": "orange"}) # 列顶部画X轴边缘直方图 g.map_upper(sns.histplot, bins=10, alpha=0.6) # 行右侧画Y轴边缘直方图 g.map_right(sns.histplot, bins=10, alpha=0.6) # 循环给每个主图标注相关系数 for row_idx, y_col in enumerate(rt): for col_idx, x_col in enumerate(nvars): r, p = stats.pearsonr(df[y_col], df[x_col]) g.axes[row_idx, col_idx].annotate(f'$\\rho = {r:.3f}, p = {p:.3f}$', xy=(0.1, 0.9), xycoords='axes fraction', ha='left', va='center', bbox={'boxstyle': 'round', 'fc': 'powderblue', 'ec': 'navy'}) plt.tight_layout() plt.show()
方案2:每个单元格带独立边缘(和单个jointplot效果一致)
如果需要每个单元格都像单独的jointplot一样自带独立边缘,需要手动在画布上创建带边缘的子图组,不调用sns.jointplot接口:
import numpy as np import pandas as pd import scipy.stats as stats import seaborn as sns import matplotlib.pyplot as plt from matplotlib.gridspec import GridSpec ncols=['ra','rb','a','b','c','d'] df=pd.DataFrame(np.random.rand(100,len(ncols)),columns=ncols) nvars=['a','b','c','d'] rt=['a','b'] n_rows = len(rt) n_cols = len(nvars) # 每个joint单元占2行2列的网格空间:主图在(1,0),X边缘在(0,0),Y边缘在(1,1) fig = plt.figure(figsize=(n_cols*4, n_rows*4)) gs = GridSpec(n_rows*2, n_cols*2, figure=fig, width_ratios=[5,1]*n_cols, height_ratios=[1,5]*n_rows, wspace=0.1, hspace=0.1) for row_idx, y_col in enumerate(rt): for col_idx, x_col in enumerate(nvars): # 计算当前单元在GridSpec中的位置 base_row = row_idx * 2 base_col = col_idx * 2 ax_main = fig.add_subplot(gs[base_row+1, base_col]) ax_xmargin = fig.add_subplot(gs[base_row, base_col], sharex=ax_main) ax_ymargin = fig.add_subplot(gs[base_row+1, base_col+1], sharey=ax_main) # 画主回归图 sns.regplot(data=df, x=x_col, y=y_col, ax=ax_main, line_kws={"color": "orange"}) # 画X边缘直方图 sns.histplot(data=df, x=x_col, ax=ax_xmargin, alpha=0.6) # 画Y边缘直方图 sns.histplot(data=df, y=y_col, ax=ax_ymargin, alpha=0.6) # 隐藏边缘图的刻度标签 ax_xmargin.tick_params(axis='x', labelbottom=False) ax_ymargin.tick_params(axis='y', labelleft=False) # 标注相关系数 r, p = stats.pearsonr(df[y_col], df[x_col]) ax_main.annotate(f'$\\rho = {r:.3f}, p = {p:.3f}$', xy=(0.1, 0.9), xycoords='axes fraction', ha='left', va='center', bbox={'boxstyle': 'round', 'fc': 'powderblue', 'ec': 'navy'}) plt.show()
内容的提问来源于stack exchange,提问作者rpb
相关产品推荐
相关产品推荐

