You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义Matplotlib Figure和Axes类时,继承Axes出现TypeError的问题求助

自定义Matplotlib Figure和Axes类时,继承Axes出现TypeError的问题求助

我尝试模仿plt.subplots()的行为,但用自定义类来实现——希望subplots()返回CustomAxes而不是默认的Axes对象。虽然我查看了Matplotlib的源码,但还是搞不懂为什么会出现下面的回溯错误。

目前我不继承Axes也能实现需求,但从长期维护的角度来看,我更希望能直接继承Axes类。如果你觉得这个思路不靠谱,有更合理的实现方式,欢迎给我提建议!

报错代码

from matplotlib.figure import Figure
from matplotlib.axes import Axes

class CustomAxes(Axes):
    
    def __init__(self, fig, *args, **kwargs):
        super().__init__(fig, *args, **kwargs)
    
    def create_plot(self, i):
        self.plot([1, 2, 3], [1, 2, 3])
        self.set_title(f'Title {i}')

class CustomFigure(Figure):

    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
    
    def subplots(self, *args, **kwargs):
        axes = super().subplots(*args, **kwargs)
        axes = [CustomAxes(fig=self, *args, **kwargs) for ax in axes.flatten()]
        return axes

fig, axes = CustomFigure().subplots(nrows=2, ncols=2)
for i, ax in enumerate(axes, start=1):
    ax.create_plot(i=i)
fig.tight_layout()

fig

报错回溯信息

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In[60], line 23
     20         axes = [CustomAxes(fig=self, *args, **kwargs) for ax in axes.flatten()]
     21         return axes
---> 23 fig, axes = CustomFigure().subplots(nrows=2, ncols=2)
     24 for i, ax in enumerate(axes, start=1):
     25     ax.create_plot(i=i)

Cell In[60], line 20
     18 def subplots(self, *args, **kwargs):
     19     axes = super().subplots(*args, **kwargs)
---> 20     axes = [CustomAxes(fig=self, *args, **kwargs) for ax in axes.flatten()]
     21     return axes

Cell In[60], line 20
     18 def subplots(self, *args, **kwargs):
     19     axes = super().subplots(*args, **kwargs)
---> 20     axes = [CustomAxes(fig=self, *args, **kwargs) for ax in axes.flatten()]
     21     return axes

Cell In[60], line 7
      6 def __init__(self, fig, *args, **kwargs):
----> 7     super().__init__(fig, *args, **kwargs)

File ~/repos/test/venv/lib/python3.11/site-packages/matplotlib/axes/_base.py:656, in _AxesBase.__init__(self, fig, facecolor, frameon, sharex, sharey, label, xscale, yscale, box_aspect, forward_navigation_events, *args, **kwargs)
    654 else:
    655     self._position = self._originalPosition = mtransforms.Bbox.unit()
--> 656     subplotspec = SubplotSpec._from_subplot_args(fig, args)
    657 if self._position.width < 0 or self._position.height < 0:
    658     raise ValueError('Width and height specified must be non-negative')

File ~/repos/test/venv/lib/python3.11/site-packages/matplotlib/gridspec.py:576, in SubplotSpec._from_subplot_args(figure, args)
    574     rows, cols, num = args
    575 else:
--> 576     raise _api.nargs_error("subplot", takes="1 or 3", given=len(args))
    578 gs = GridSpec._check_gridspec_exists(figure, rows, cols)
    579 if gs is None:

TypeError: subplot() takes 1 or 3 positional arguments but 0 were given

不继承Axes的可行代码

from matplotlib.figure import Figure


class CustomAxes():

    def __init__(self, ax):
        self.ax = ax

    def create_plot(self, i):
        self.ax.plot([1, 2, 3], [1, 2, 3])
        self.ax.set_title(f'Title {i}')


class CustomFigure(Figure):

    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)

    def subplots(self, *args, **kwargs):
        axes = super().subplots(*args, **kwargs)
        axes = [CustomAxes(ax) for ax in axes.flatten()]
        return self, axes


fig, axes = CustomFigure().subplots(nrows=2, ncols=2)
for i, ax in enumerate(axes, start=1):
    ax.create_plot(i=i)
fig.tight_layout()

fig

备注:内容来源于stack exchange,提问作者Simon1

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 14:53:12