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

Matplotlib/Seaborn中3×3热力图的子图独立高度设置

3×3热力图子图独立高度设置方案

问题描述

需要创建3×3布局的热力图,要求:

  • 所有9张子图宽度一致
  • 每个子图高度按「热力图单元格数量×固定单元格高度」设置,实现所有热力图单元格物理高度统一(例如第一行第二列的子图仅1行单元格,高度对应1个单元格)
    当前代码仅能调整整行子图高度,导致同行内所有子图高度一致,无法满足需求。当前效果如下:

当前热力图效果

解决方案

核心思路是手动计算每个子图的位置和尺寸,替代原有的整行GridSpec布局,实现每个子图独立设置高度,确保单元格高度统一。具体步骤:

  1. 计算每个分组的单元格行数,对应子图高度
  2. 按行确定子图的顶部对齐位置,保证同行子图顶部在同一水平线
  3. 手动创建每个子图的axes,设置对应的宽度、高度和位置
  4. 保留原有热力图绘制逻辑,适配新的axes布局

修改后的完整代码

import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns

groups = ['g1', 'g2', 'g3', 'g4', 'g5', 'g6', 'g7', 'g8', 'g9']

dummy_data = {
    'names':['name_1', 'name_2', 'name_3', 'name_4', 'name_5', 'name_6', 'name_7', 'name_8', 'name_9', 'name_10', 'name_11', 'name_12', 'name_13', 'name_14', 'name_15', 'name_16', 'name_17', 'name_18', 'name_19', 'name_20'],
    'group': [
        'g1', 'g1', 'g1', 'g1', 'g2', 'g3', 'g3', 'g3', 'g4', 'g4',
        'g5', 'g5', 'g6', 'g6', 'g7', 'g7', 'g8', 'g9', 'g9', 'g9'
    ],
    'col1': [
        -30, 10, 5, -20, 15, 25, -15, 0, -10, 20,
        -5, 15, 35, -25, 10, 5, -15, 30, 10, -5
    ],
    'col2': [
        -50, 20, 30, -40, 25, 45, -10, 5, -15, 35,
        -10, 20, 40, -20, 15, 10, -20, 45, 20, -10
    ],
    'col3': [
        -45, 15, 35, -35, 20, 40, -5, 10, -10, 30,
        -15, 25, 45, -15, 12, 7, -18, 40, 25, -8
    ],
    'col4': [
        0.05, 0.08, 0.07, 0.06, 0.09, 0.11, 0.05, 0.08, 0.07, 0.10,
        0.04, 0.09, 0.12, 0.06, 0.07, 0.08, 0.06, 0.11, 0.05, 0.09
    ],
    'col5': [
        0.12, 0.07, 0.09, 0.04, 0.08, 0.10, 0.06, 0.09, 0.05, 0.11,
        0.03, 0.10, 0.13, 0.07, 0.06, 0.07, 0.05, 0.10, 0.04, 0.08
    ],
    'col6': [
        20, 25, 15, 30, 40, 10, 35, 5, 45, 25,
        10, 20, 50, 15, 30, 25, 40, 15, 10, 35
    ],
    'col7': [
        45, 15, 25, 10, 30, 40, 20, 35, 5, 20,
        30, 15, 40, 20, 10, 15, 50, 30, 5, 25
    ],
    'col8': [
        5, 2, 8, 3, 7, 6, 4, 9, 1, 8,
        6, 3, 10, 2, 7, 4, 9, 5, 2, 8
    ],
    'col9': [
        6, 3, 9, 4, 8, 7, 5, 10, 2, 9,
        7, 4, 10, 3, 8, 5, 9, 6, 3, 9
    ]
}

pivot_df = pd.DataFrame(dummy_data)

# Define column bounds with new names
column_bounds = {
    'col1': (-35, 35),
    'col2': (-55, 55),
    'col3': (-50, 50),
    'col4': (0.03, 0.15),
    'col5': (0.03, 0.15),
    'col6': (0, 50),
    'col7': (0, 50),
    'col8': (0, 10),
    'col9': (0, 10),
}

# Calculate row counts for each group
row_counts = [len(pivot_df.loc[pivot_df['group'] == group]) for group in groups]

# Set fixed cell height and calculate each subplot's height
cell_height = 0.5  # Fixed height per cell
subplot_heights = [count * cell_height for count in row_counts]

# Layout parameters
fig_width = 32
wspace = 0.4  # Horizontal space between subplots
hspace = 0.2  # Vertical space between rows

# Split subplots into 3 rows
row_subplots = [[0, 1, 2], [3, 4, 5], [6, 7, 8]]
# Calculate max height for each row (to align tops)
row_max_heights = [max(subplot_heights[i] for i in row) for row in row_subplots]
# Total figure height including spaces
total_height = sum(row_max_heights) + hspace * (len(row_subplots) - 1)

# Create figure
fig = plt.figure(figsize=(fig_width, total_height))
axes = []

current_top = 1.0  # Start from top of figure (0-1 coordinate system)

for row in row_subplots:
    row_max_h = max(subplot_heights[i] for i in row)
    # Calculate width per subplot (accounting for horizontal space)
    subplot_width = (1 - wspace * (len(row) - 1)) / len(row)
    
    for idx, subplot_idx in enumerate(row):
        # Calculate left position of current subplot
        left = idx * (subplot_width + wspace)
        # Height of current subplot in figure coordinates
        h = subplot_heights[subplot_idx] / total_height
        # Bottom position (align top with row's top)
        bottom = current_top - h
        
        # Add subplot to figure
        ax = fig.add_axes([left, bottom, subplot_width, h])
        axes.append(ax)
    
    # Move current top down to next row's top (accounting for vertical space)
    current_top -= row_max_h / total_height + hspace / total_height

# Plot heatmaps for each group
for i, group in enumerate(groups):
    # Get the data for the current group
    group_data = pivot_df.loc[pivot_df['group'] == group].set_index('names').drop(columns={'group'})

    # Create annotation array
    annotations = group_data.copy()
    for col in ['col6', 'col7', 'col8', 'col9']:
        annotations[col] = annotations[col].apply(lambda x: f"{x:.2f}M")
    for col in ['col1', 'col2', 'col3', 'col4', 'col5']:
        annotations[col] = annotations[col].apply(lambda x: f"{x:.2f}")
    
    # Plot each column individually with separate bounds
    for j, col in enumerate(group_data.columns):
        # Create a mask for other columns to plot each individually
        mask = np.ones(group_data.shape)
        mask[:, j] = 0  # Mask all except the current column
       
        sns.heatmap(
            group_data,
            annot=annotations,
            fmt='',
            cmap="coolwarm",
            cbar=False,
            vmin=column_bounds[col][0], 
            vmax=column_bounds[col][1],
            mask=mask, 
            ax=axes[i]
        )

    axes[i].set_title(group)
    axes[i].set_ylabel('')

    # Rotate y-tick labels to be horizontal if there are any labels
    yticklabels = group_data.index.tolist()
    axes[i].set_yticks(np.arange(len(group_data)))  # Set tick positions based on number of rows
    axes[i].set_yticklabels(yticklabels, rotation=0, ha='right', fontsize=8)

    if i < 6:  # Only for the top two rows (0–5), remove x labels and ticks
        axes[i].set_xlabel('')
        axes[i].set_xticks([])

plt.show()

关键改动说明

  1. 子图位置计算:放弃原有的GridSpec整行布局,改用fig.add_axes()手动定义每个子图的位置和尺寸,确保每个子图高度由对应分组的单元格数量决定
  2. 顶部对齐逻辑:按行计算最大子图高度,使同行内所有子图顶部在同一水平线,保证布局整齐
  3. 间距适配:通过wspace和hspace参数控制子图间的水平、垂直间距,保持布局美观

内容的提问来源于stack exchange,提问作者L M

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 10:10:56