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

Matplotlib中逐步添加图例自定义条目时的覆盖问题及解决思路

Matplotlib分步构建图例及分组信息传达方案

问题重现

你当前的代码尝试通过plot_on_axis函数分步向轴添加绘图元素并更新图例,但第二次调用时会丢失第一次添加的分组patch(sin条目),原因是ax.get_legend_handles_labels()仅返回当前轴上的绘图线条,不会包含之前手动添加的patch,导致图例被覆盖。

示例代码:

import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
import numpy as np

def plot_on_axis(ax: plt.Axes, x: np.ndarray, y: np.ndarray, color, name) -> plt.Axes:
    ax.plot(x, y, color=color, label="orig")
    ax.plot(x, y + 0.2, "--", color=color, label="shifted")
    patch = mpatches.Patch(color=color, label=name)
    handles, labels = ax.get_legend_handles_labels()
    ax.legend(handles + [patch], labels + [name])
    return ax

def get_fig() -> plt.Figure:
    x1 = np.linspace(0, 3)
    y1 = np.sin(x1)

    x2 = np.linspace(0, 3)
    y2 = np.cos(x2)

    fig = plt.figure()
    ax = fig.subplots()
    plot_on_axis(ax, x1, y1, "tab:blue", "sin")
    plot_on_axis(ax, x2, y2, "tab:orange", "cos")

    return fig

get_fig().show()

解决方案

1. 在plot_on_axis中逐步构建图例

完全可以在plot_on_axis函数内实现图例的逐步累加,不需要移到get_fig中处理。核心思路是自己维护图例的handles和labels集合,而不是依赖ax.get_legend_handles_labels()的返回值。

修改后的代码:

import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
import numpy as np

def plot_on_axis(ax: plt.Axes, x: np.ndarray, y: np.ndarray, color, name) -> plt.Axes:
    # 绘制当前组的两条曲线,保存返回的线条对象
    line_orig, = ax.plot(x, y, color=color, label="orig")
    line_shifted, = ax.plot(x, y + 0.2, "--", color=color, label="shifted")
    # 创建分组标识patch
    group_patch = mpatches.Patch(color=color, label=name)
    
    # 初始化轴对象的自定义属性,用于存储累计的图例条目
    if not hasattr(ax, '_legend_handles'):
        ax._legend_handles = []
        ax._legend_labels = []
    
    # 将当前组的条目加入累计集合
    ax._legend_handles.extend([line_orig, line_shifted, group_patch])
    ax._legend_labels.extend(["orig", "shifted", name])
    
    # 更新图例
    ax.legend(ax._legend_handles, ax._legend_labels)
    return ax

def get_fig() -> plt.Figure:
    x1 = np.linspace(0, 3)
    y1 = np.sin(x1)

    x2 = np.linspace(0, 3)
    y2 = np.cos(x2)

    fig = plt.figure()
    ax = fig.subplots()
    plot_on_axis(ax, x1, y1, "tab:blue", "sin")
    plot_on_axis(ax, x2, y2, "tab:orange", "cos")

    return fig

get_fig().show()

原理说明

  • 给轴对象添加自定义属性_legend_handles和_legend_labels,用于存储所有累计的图例条目。
  • 每次调用plot_on_axis时,将当前组的线条和分组patch加入集合,再用这个集合更新图例,避免了仅从轴上获取绘图元素导致的条目丢失。

2. 更清晰的分组信息传达方式

除了添加分组patch,还有几种更直观的方式来区分不同数据组:

方式一:嵌套图例

创建主图例显示分组名称,子图例显示具体条目,结构层级更清晰:

import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
import numpy as np

def plot_on_axis(ax: plt.Axes, x: np.ndarray, y: np.ndarray, color, name) -> plt.Axes:
    ax.plot(x, y, color=color, label=f"{name}: orig")
    ax.plot(x, y + 0.2, "--", color=color, label=f"{name}: shifted")
    return ax

def get_fig() -> plt.Figure:
    x1 = np.linspace(0, 3)
    y1 = np.sin(x1)

    x2 = np.linspace(0, 3)
    y2 = np.cos(x2)

    fig = plt.figure()
    ax = fig.subplots()
    plot_on_axis(ax, x1, y1, "tab:blue", "sin")
    plot_on_axis(ax, x2, y2, "tab:orange", "cos")

    # 获取所有绘图条目
    all_handles, all_labels = ax.get_legend_handles_labels()
    # 拆分sin和cos的条目
    sin_handles, sin_labels = all_handles[:2], all_labels[:2]
    cos_handles, cos_labels = all_handles[2:], all_labels[2:]

    # 主图例:显示分组名称
    sin_patch = mpatches.Patch(color="tab:blue", label="sin")
    cos_patch = mpatches.Patch(color="tab:orange", label="cos")
    main_legend = ax.legend([sin_patch, cos_patch], ["sin", "cos"], loc="upper left")
    ax.add_artist(main_legend)  # 保留主图例,避免被覆盖

    # 子图例:显示具体条目
    ax.legend(sin_handles + cos_handles, sin_labels + cos_labels, loc="lower right")

    return fig

get_fig().show()

方式二:带标题的分组图例

通过添加无颜色的补丁作为分组标题,结合多列布局让分组更直观:

import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
import numpy as np

def plot_on_axis(ax: plt.Axes, x: np.ndarray, y: np.ndarray, color, name) -> plt.Axes:
    ax.plot(x, y, color=color, label=f"{name}: orig")
    ax.plot(x, y + 0.2, "--", color=color, label=f"{name}: shifted")
    return ax

def get_fig() -> plt.Figure:
    x1 = np.linspace(0, 3)
    y1 = np.sin(x1)

    x2 = np.linspace(0, 3)
    y2 = np.cos(x2)

    fig = plt.figure()
    ax = fig.subplots()
    plot_on_axis(ax, x1, y1, "tab:blue", "sin")
    plot_on_axis(ax, x2, y2, "tab:orange", "cos")

    all_handles, all_labels = ax.get_legend_handles_labels()
    # 添加分组标题补丁(无颜色,仅作为文字分隔)
    sin_title = mpatches.Patch(color="none", label="=== sin 组 ===")
    cos_title = mpatches.Patch(color="none", label="=== cos 组 ===")

    # 重新排列条目顺序:标题+对应组条目
    ordered_handles = [sin_title] + all_handles[:2] + [cos_title] + all_handles[2:]
    ordered_labels = ["=== sin 组 ==="] + all_labels[:2] + ["=== cos 组 ==="] + all_labels[2:]

    # 用2列布局让分组更紧凑
    ax.legend(ordered_handles, ordered_labels, ncol=2, loc="upper center")

    return fig

get_fig().show()

方式三:统一组样式+标签前缀

给同组的所有条目添加统一的颜色和样式前缀,让查看者快速识别分组:
比如将标签设为sin: orig、sin: shifted,配合统一的蓝色,无需额外添加分组patch,也能清晰区分。


内容的提问来源于stack exchange,提问作者VaNa

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 04:07:11