Matplotlib/Seaborn中3×3热力图的子图独立高度设置
3×3热力图子图独立高度设置方案
问题描述
需要创建3×3布局的热力图,要求:
- 所有9张子图宽度一致
- 每个子图高度按「热力图单元格数量×固定单元格高度」设置,实现所有热力图单元格物理高度统一(例如第一行第二列的子图仅1行单元格,高度对应1个单元格)
当前代码仅能调整整行子图高度,导致同行内所有子图高度一致,无法满足需求。当前效果如下:

解决方案
核心思路是手动计算每个子图的位置和尺寸,替代原有的整行GridSpec布局,实现每个子图独立设置高度,确保单元格高度统一。具体步骤:
- 计算每个分组的单元格行数,对应子图高度
- 按行确定子图的顶部对齐位置,保证同行子图顶部在同一水平线
- 手动创建每个子图的axes,设置对应的宽度、高度和位置
- 保留原有热力图绘制逻辑,适配新的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()
关键改动说明
- 子图位置计算:放弃原有的
GridSpec整行布局,改用fig.add_axes()手动定义每个子图的位置和尺寸,确保每个子图高度由对应分组的单元格数量决定 - 顶部对齐逻辑:按行计算最大子图高度,使同行内所有子图顶部在同一水平线,保证布局整齐
- 间距适配:通过
wspace和hspace参数控制子图间的水平、垂直间距,保持布局美观
内容的提问来源于stack exchange,提问作者L M
相关产品推荐
相关产品推荐

