🧑‍🎨

matplotlibで軸を共有した複数プロットを作成する

に公開

発端

こういう画像が作りたかった訳です。
Planck衛星によるCMB温度異方性のパワースペクトル(上)と理論値との差(下) Planck Collaboration, A&A 594, A11 (2016)
画像は https://www.cosmos.esa.int/web/planck/picture-gallery より。

matplotlibで使われる用語

本題に入る前に,matplotlibで使われる用語をいくつか振り返ります。
以下の図が分かりやすいのでお借りします。
https://matplotlib.org/1.5.1/faq/usage_faq.html より引用[1]

  • figure:作成する画像の大枠となる部分
  • axes:実際に図が置かれる部分
    • (最低限x軸とy軸があるので"axes"と理解しています)
  • axis:axesの周りにある枠
    という理解をしています。

何をしたいのか

普段はfigure.add_subplot()でaxesを作成したり,plt.subplots()でfigureとaxesを同時に作成していると思います。しかし,このときaxesの間には空間が生まれてしまいます。axes毎に違うパラメータを表示しているなら問題ありませんが,すべてのaxesでx軸,あるいはy軸のパラメータと範囲が共通なら個別に表示するだけスペースの無駄となってしまいます。余分な空間を減らして見やすさを向上するために,axisを共有したaxesを作りたいと思ったのです。

どうしたか

matplotlib.pyplotには,add_subplotよりも低レイヤー?なaxes作成関数としてadd_axesがあります。この関数の必須の引数rectは4つの値(left, bottom, width, height)を取ります。これはfigureの左下を0, 右上を1として,axesの左端の値,下端の値,横幅,縦幅の指定となります。つまり,これを上手く調整して上げれば,好きに軸の重なったプロットが作成可能です。[2]
この思想のもと作成した関数がこちらです。

def share_plot(fig: plt.Figure,
               row: int, col: int,
               rect=[0.15, 0.12, 0.10, 0.05]) -> List[plt.Axes]:
    '''
    make axes that sharing x-axis and y-axis.

    Parameters
    ----------
    fig: matplotlib.pyplot.Figure
        input figure to add axes.
    row, col: int
        row and collum number of axes
    rect: sequence of float
        margin of new axes, [left, bottom, right, top].

    Returns
    -------
    axes: Axes
        sequence of axes
    '''
    width = (1-rect[0]-rect[2])/col
    height = (1-rect[1]-rect[3])/row
    axes = []
    for i in range(row):
        for j in range(col):
            left = rect[0]+j*width
            bot = rect[1]+(row-i-1)*height
            axis = fig.add_axes((left, bot, width, height))
            axes.append(axis)
    return axes

https://github.com/MeF0504/pymeflib

引数のfigはFigureオブジェクト,rowcolは作られるaxesの行数と列数を指定します。rectadd_axesに似ていますが,幅ではなく全axesの端の値(左,下,右,上)となっています。なお,デフォルト値はmatplotlibの初期値にしています。
処理としては単純で,rectの値から各axesの大きさを計算し,for分で順番に隙間がないように配置しています。
こちらを使って,例えば

#! /usr/bin/env python3

import numpy as np
import matplotlib.pyplot as plt

from pymeflib import plot as mefplot


def main():
    x = np.arange(-10, 10, 0.01)
    y1 = x**2
    y2 = x**3
    fig1 = plt.figure()
    ax11, ax12 = mefplot.share_plot(fig1, 2, 1)
    ax11.plot(x, y1, 'r', label=r'$x^2$')
    ax12.plot(x, y2, 'b', label=r'$x^3$')
    ax11.set_xticklabels([])
    ax12.set_xlabel(r'$x$')
    ax11.set_ylabel(r'$x^2$')
    ax12.set_ylabel(r'$x^3$')
    ax11.legend()
    ax12.legend()
    plt.show()


if __name__ == '__main__':
    main()

とすると,以下のような図が作れます。

注意として,軸の範囲合わせや軸ラベルの除去はしていないので,適宜行ってください。

終わりに

今回は自分の作成した関数の紹介でした。add_axesはレイヤーが低い?ため設定内容が分かりにくいですが,慣れると今回のような隙間なくaxesを配置したり,大きさのバラバラなaxesを配置したり,大きなaxesの中に小さいaxesを配置したり出来るので,この記事をきっかけに触ってもらえると嬉しいです。

脚注
  1. 画像を知ったきっかけはこちら ↩︎

  2. subplots_adjust というsubplotsの間隔設定用の関数があるんですね.. 参考 ↩︎

Discussion