解决SeabornFig2Grid中jointplot色条重叠及标签截断问题
解决方案
核心思路是不在jointplot的独立画布中创建颜色条,而是在最终的大画布中统一分配颜色条位置,同时优化网格布局比例,避免元素重叠和标签截断:
关键修改点
- 保存
hist2d返回的图像对象,用于后续在大画布中创建颜色条 - 删除jointplot画布中创建颜色条的代码,避免位置冲突
- 调整GridSpec的宽度比例,给颜色条预留足够空间,同时适配整体画布尺寸
- 在大画布中为每个jointplot单独创建颜色条,并通过
gridspec绑定对应位置
修改后的完整代码
import matplotlib import seaborn as sns import numpy as np import matplotlib.pyplot as plt import matplotlib.colors as mcolors import matplotlib.gridspec as gridspec class SeabornFig2Grid(): """Allow seaborn figure-level figs to be suplots.""" def __init__(self, seaborngrid, fig, subplot_spec): self.fig = fig self.sg = seaborngrid self.subplot = subplot_spec if isinstance(self.sg, sns.axisgrid.FacetGrid) or \ isinstance(self.sg, sns.axisgrid.PairGrid): self._movegrid() elif isinstance(self.sg, sns.axisgrid.JointGrid): self._movejointgrid() elif isinstance(self.sg, matplotlib.axes._axes.Axes): self._moveaxis() self._finalize() def _moveaxis(self): self._resize() self.subgrid = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=self.subplot) self._moveaxes(self.sg, self.subgrid[0, 0]) def _movegrid(self): """Move PairGrid or Facetgrid.""" self._resize() n = self.sg.axes.shape[0] m = self.sg.axes.shape[1] self.subgrid = gridspec.GridSpecFromSubplotSpec(n, m, subplot_spec=self.subplot) for i in range(n): for j in range(m): self._moveaxes(self.sg.axes[i, j], self.subgrid[i, j]) def _movejointgrid(self): """Move Jointgrid.""" h = self.sg.ax_joint.get_position().height h2 = self.sg.ax_marg_x.get_position().height r = int(np.round(h / h2)) self._resize() self.subgrid = gridspec.GridSpecFromSubplotSpec(r + 1, r + 1, subplot_spec=self.subplot) self._moveaxes(self.sg.ax_joint, self.subgrid[1:, :-1]) self._moveaxes(self.sg.ax_marg_x, self.subgrid[0, :-1]) self._moveaxes(self.sg.ax_marg_y, self.subgrid[1:, -1]) def _moveaxes(self, ax, grid_spec): ax.remove() ax.figure = self.fig self.fig.axes.append(ax) self.fig.add_axes(ax) ax._subplotspec = grid_spec ax.set_position(grid_spec.get_position(self.fig)) try: ax.set_subplotspec(grid_spec) except AttributeError: ax._subplotspec = grid_spec def _finalize(self): try: plt.close(self.sg.fig) except AttributeError: pass self.fig.canvas.mpl_connect("resize_event", self._resize) self.fig.canvas.draw() def _resize(self, evt=None): self.sg.figure.set_size_inches(self.fig.get_size_inches()) # 生成数据 x = [np.random.random() for _ in range(1000)] y = [np.random.random() for _ in range(1000)] # 创建jointplot并保存图像对象 joint_plots = [] hist_images = [] for _ in range(2): # 创建jointplot jp = sns.jointplot(x=x, y=y, marginal_kws={'bins': 20}) jp.ax_joint.cla() # 在joint轴绘制2D直方图并保存图像对象 plt.sca(jp.ax_joint) im = plt.hist2d(x, y, bins=20, norm=mcolors.LogNorm(), cmap='jet') hist_images.append(im[3]) # 获取图像对象 joint_plots.append(jp) # 设置整体画布和网格布局 fig = plt.figure(figsize=(10, 4)) # 1行4列:两个jointplot区域 + 两个颜色条区域,宽度比例适配 gs = gridspec.GridSpec(1, 4, width_ratios=[4, 0.2, 4, 0.2]) # 将jointplot添加到布局中 sfg_list = [] for idx, jp in enumerate(joint_plots): sfg = SeabornFig2Grid(jp, fig, gs[0, idx*2]) sfg_list.append(sfg) # 在大画布中添加颜色条 for idx, im in enumerate(hist_images): cbar_ax = fig.add_subplot(gs[0, idx*2+1]) cb = plt.colorbar(im, cax=cbar_ax) cb.set_label(r"$\log_{10}$ density of points", fontsize=13) # 调整整体布局,避免边缘截断 plt.subplots_adjust(left=0.05, right=0.95, top=0.9, bottom=0.1) plt.show()
内容的提问来源于stack exchange,提问作者BML
相关产品推荐
相关产品推荐

