如何在Matplotlib的plt.subplots()中绘制多个边缘KDE散点组合图
实现4个散点图+边缘KDE组合子图
问题背景
已通过Matplotlib结合Seaborn实现单张包含散点图与边缘KDE曲线的组合图,现需在plt.subplots()体系下绘制4个该类结构的子图。
解决方案思路
把单个组合图的绘制逻辑封装成可复用函数,通过嵌套网格布局(GridSpec)在大画布上划分4个独立区域,每个区域内部构建2x2的子轴网格,完成散点图+边缘KDE的组合结构绘制。
完整代码示例
import matplotlib.pyplot as plt import seaborn as sns import pandas as pd import numpy as np # 模拟4组待可视化数据,替换为你实际的4组results数据 data_sets = [ pd.DataFrame({'real': np.random.normal(0, 1, 1000), 'pred': np.random.normal(0, 1, 1000)}), pd.DataFrame({'real': np.random.normal(1, 1.2, 1000), 'pred': np.random.normal(1, 1.2, 1000)}), pd.DataFrame({'real': np.random.normal(-1, 0.8, 1000), 'pred': np.random.normal(-1, 0.8, 1000)}), pd.DataFrame({'real': np.random.normal(2, 1.5, 1000), 'pred': np.random.normal(2, 1.5, 1000)}) ] def draw_joint_plot(gs, data, xlabel, ylabel): # 在指定主网格位置创建2x2子轴网格 axs_sub = plt.subplot(gs).subgridspec(2, 2, hspace=0, wspace=0, width_ratios=[5, 1], height_ratios=[1, 5]) ax00 = plt.subplot(axs_sub[0, 0]) ax01 = plt.subplot(axs_sub[0, 1]) ax10 = plt.subplot(axs_sub[1, 0]) ax11 = plt.subplot(axs_sub[1, 1]) # 关闭冗余轴 ax00.axis("off") ax01.axis("off") ax11.axis("off") # 绘制边缘KDE曲线 sns.kdeplot(data['real'], fill=True, ax=ax00) sns.kdeplot(y=data['pred'], fill=True, ax=ax11) # 绘制散点图与参考线 ax10.scatter(data['real'], data['pred'], s=2) ax10.axline([0, 0], [1, 1], linestyle='--', color='red') # 设置轴标签 ax10.set_xlabel(xlabel, labelpad=10) ax10.set_ylabel(ylabel, labelpad=10) # 设置内侧刻度 ax10.tick_params(top=True, right=True, direction='in') # 同步轴限消除空隙 ax00.set_xlim(ax10.get_xlim()) ax11.set_ylim(ax10.get_ylim()) # 创建大画布与主网格布局 fig = plt.figure(figsize=(12, 10)) main_gs = fig.add_gridspec(2, 2, hspace=0.3, wspace=0.3) # 批量绘制4个组合图 draw_joint_plot(main_gs[0,0], data_sets[0], 'Real $\log_{10}(\sigma)$', 'Pred. $\log_{10}(\sigma)$') draw_joint_plot(main_gs[0,1], data_sets[1], 'Real $\log_{10}(\sigma)$', 'Pred. $\log_{10}(\sigma)$') draw_joint_plot(main_gs[1,0], data_sets[2], 'Real $\log_{10}(\sigma)$', 'Pred. $\log_{10}(\sigma)$') draw_joint_plot(main_gs[1,1], data_sets[3], 'Real $\log_{10}(\sigma)$', 'Pred. $\log_{10}(\sigma)$') plt.tight_layout() plt.show()
代码说明
- 复用函数设计:
draw_joint_plot函数接收网格位置、数据和标签参数,封装单个组合图的全部绘制逻辑,避免重复代码。 - 嵌套网格布局:通过
subgridspec在每个主网格区域内构建2x2子网格,精准实现散点图与边缘KDE的排布。 - 视觉一致性:统一设置刻度样式、轴限同步规则,保证4个子图的视觉风格协调统一。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

