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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 01:00:20