matplotlibで軸を共有した複数プロットを作成する
発端
こういう画像が作りたかった訳です。
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
引数のfigはFigureオブジェクト,rowとcolは作られるaxesの行数と列数を指定します。rectはadd_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を配置したり出来るので,この記事をきっかけに触ってもらえると嬉しいです。
Discussion