如何在Matplotlib子图中让1:1参考线居中且保留独立轴限?
问题描述
我创建了一个2x2的Matplotlib子图网格,每个子图包含不同数据点的散点图。我尝试在每个子图中绘制斜率为1、截距为0的abline(1:1参考线)以可视化数据关系,但因各子图数据范围不同,参考线无法在所有子图中保持1:1居中。
我希望在保留每个子图基于自身数据的独立轴限的前提下,让参考线在每个子图中居中,即参考线穿过子图数据点的中心且不扭曲数据。
当前代码如下:
timesteps = [185, 159, 53, 2] def abline(ax, slope, intercept): """Plot a line from slope and intercept""" x_vals = np.array(ax.get_xlim()) y_vals = intercept + slope * x_vals ax.plot(x_vals, y_vals, 'r--') fig, axs = plt.subplots(2, 2, figsize=(12, 8)) for i, timestep in enumerate(timesteps): mask = np.where(nan_mask[timestep, :, :] == 0) data_tmwm_values = data_tmwm[timestep, :, :][mask] ds_plot_values = ds_og_red[timestep, :, :][mask] row = i // 2 # Integer division to get the row index col = i % 2 # Modulo operation to get the column index ax = axs[row, col] ax.scatter(data_tmwm_values, ds_plot_values, s=20) ax.set_xlabel('TMWM') ax.set_ylabel('Original') ax.set_title(f'Scatter Plot (Timestep: {timestep})') correlation_matrix = np.corrcoef(data_tmwm_values, ds_plot_values) r_value = correlation_matrix[0, 1] r_squared = r_value ** 2 abline(ax, 1, 0) ax.text(0.05, 0.95, f"R² value: {r_squared:.3f}", transform=ax.transAxes, ha='left', va='top') plt.tight_layout() plt.show()
当前效果:
我已尝试使用get_xlim()和get_ylim()函数设置轴限,但未能实现参考线的正确居中,恳请指导如何达成需求。
解决方案
问题核心在于:当前参考线基于子图自动生成的轴限绘制,当x、y轴范围不一致时,1:1线会视觉上偏移。要解决这个问题,需让每个子图的x、y轴采用相同的轴限范围(覆盖当前子图所有数据),这样1:1线就能自然居中,且不扭曲数据展示。
具体修改步骤
统一每个子图的x、y轴限
在绘制散点图后,计算当前子图所有数据的极值,将x、y轴范围设置为该极值区间,确保轴范围一致:# 在ax.scatter(...)之后添加 all_data = np.concatenate([data_tmwm_values, ds_plot_values]) data_min = all_data.min() data_max = all_data.max() # 设置x、y轴范围完全一致 ax.set_xlim(data_min, data_max) ax.set_ylim(data_min, data_max)优化abline函数(可选)
让参考线覆盖整个轴范围,避免因轴限调整后线的长度不足:def abline(ax, slope, intercept): """Plot a 1:1 line covering the full axis range""" plot_min, plot_max = ax.get_xlim() x_vals = np.array([plot_min, plot_max]) y_vals = intercept + slope * x_vals ax.plot(x_vals, y_vals, 'r--')完整修改后的代码
import numpy as np import matplotlib.pyplot as plt timesteps = [185, 159, 53, 2] def abline(ax, slope, intercept): """Plot a 1:1 line covering the full axis range""" plot_min, plot_max = ax.get_xlim() x_vals = np.array([plot_min, plot_max]) y_vals = intercept + slope * x_vals ax.plot(x_vals, y_vals, 'r--') fig, axs = plt.subplots(2, 2, figsize=(12, 8)) for i, timestep in enumerate(timesteps): mask = np.where(nan_mask[timestep, :, :] == 0) data_tmwm_values = data_tmwm[timestep, :, :][mask] ds_plot_values = ds_og_red[timestep, :, :][mask] row = i // 2 col = i % 2 ax = axs[row, col] ax.scatter(data_tmwm_values, ds_plot_values, s=20) ax.set_xlabel('TMWM') ax.set_ylabel('Original') ax.set_title(f'Scatter Plot (Timestep: {timestep})') # 统一x、y轴范围为当前子图数据的极值 all_data = np.concatenate([data_tmwm_values, ds_plot_values]) data_min = all_data.min() data_max = all_data.max() ax.set_xlim(data_min, data_max) ax.set_ylim(data_min, data_max) correlation_matrix = np.corrcoef(data_tmwm_values, ds_plot_values) r_value = correlation_matrix[0, 1] r_squared = r_value ** 2 abline(ax, 1, 0) ax.text(0.05, 0.95, f"R² value: {r_squared:.3f}", transform=ax.transAxes, ha='left', va='top') plt.tight_layout() plt.show()额外优化(可选)
如果想让数据点和轴边缘留有间隙,可添加padding:padding = (data_max - data_min) * 0.05 ax.set_xlim(data_min - padding, data_max + padding) ax.set_ylim(data_min - padding, data_max + padding)
修改后,每个子图的x、y轴范围完全匹配,1:1参考线会穿过数据中心且视觉居中,同时每个子图保留自身数据的独立轴限,不会互相干扰。
内容的提问来源于stack exchange,提问作者LouiT
相关产品推荐
相关产品推荐

