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

如何修正GluonTS绘制的Matplotlib图表中图例颜色错误问题

问题:GluonTS预测图表置信区间图例颜色异常&图例获取失败

问题背景

我使用GluonTS绘制预测图表,代码来自DeepVaR notebook,当前遇到两个核心问题:

  • 图例中置信区间的颜色显示不正确
  • 调用plt.gca().get_legend_handles_labels()仅能获取到观测值的线条,无法获取预测相关的图例项

我的绘图代码:

def plot_prob_forecasts(ts_entry, forecast_entry, asset_name, plot_length=20):
    prediction_intervals = (0.95, 0.99)
    legend = ["observations", "median prediction"] + [f"{k}% prediction interval" for k in prediction_intervals][::-1]
    fig, ax = plt.subplots(1, 1, figsize=(10, 7))
    ts_entry[-plot_length:].plot(ax=ax)  # plot the time series
    forecast_entry.plot( intervals=prediction_intervals, color='g')    
    plt.grid(which="both")
    plt.legend(legend, loc="upper left")
    plt.title(f'Forecast of {asset_name} series Returns')
    plt.show()

GluonTS内置的plot方法代码:

def plot(
    self,
    *,
    intervals=(0.5, 0.9),
    ax=None,
    color=None,
    name=None,
    show_label=False,
):
    """
    Plot median forecast and prediction intervals using ``matplotlib``.

    By default the `0.5` and `0.9` prediction intervals are plotted. Other
    intervals can be choosen by setting `intervals`.

    This plots to the current axes object (via ``plt.gca()``), or to ``ax``
    if provided. Similarly, the color is using matplotlibs internal color
    cycle, if no explicit ``color`` is set.

    One can set ``name`` to use it as the ``label`` for the median
    forecast. Intervals are not labeled, unless ``show_label`` is set to
    ``True``.
    """
    import matplotlib.pyplot as plt

    # Get current axes (gca), if not provided explicitly.
    ax = maybe.unwrap_or_else(ax, plt.gca)

    # If no color is provided, we use matplotlib's internal color cycle.
    # Note: This is an internal API and might change in the future.
    color = maybe.unwrap_or_else(
        color, lambda: ax._get_lines.get_next_color()
    )

    # Plot median forecast
    ax.plot(
        self.index.to_timestamp(),
        self.quantile(0.5),
        color=color,
        label=name,
    )

    # Plot prediction intervals
    for interval in intervals:
        if show_label:
            if name is not None:
                label = f"{name}: {interval}"
            else:
                label = interval
        else:
            label = None

        # Translate interval to low and high values. E.g for `0.9` we get
        # `low = 0.05` and `high = 0.95`. (`interval + low + high == 1.0`)
        # Also, higher interval values mean lower confidence, and thus we
        # we use lower alpha values for them.
        low = (1 - interval) / 2
        ax.fill_between(
            # TODO: `index` currently uses `pandas.Period`, but we need
            # to pass a timestamp value to matplotlib. In the future this
            # will use ``zebras.Periods`` and thus needs to be adapted.
            self.index.to_timestamp(),
            self.quantile(low),
            self.quantile(1 - low),
            # Clamp alpha betwen ~16% and 50%.
            alpha=0.5 - interval / 3,
            facecolor=color,
            label=label,
        )

环境版本:

  • python=3.9.18
  • matplotlib=3.8.0
  • gluonts=0.13.2

已尝试但无效的操作:

  • 设置color=None触发Matplotlib报错
  • 设置show_label=True并传入name无法解决图例问题

问题分析

  1. 图例颜色异常:手动指定legend列表时,Matplotlib无法将自定义标签与fill_between生成的区间图形关联,导致颜色匹配混乱。
  2. 图例项无法获取:默认show_label=False时,fill_between的label被设为None,且ax.plot的name未传入,导致预测相关元素未被添加到图例句柄集合中。

修复方法

方法1:利用GluonTS参数自动生成匹配图例

修改绘图代码,传入name和show_label=True,让GluonTS自动为中位数和区间生成标签,无需手动构造legend列表:

def plot_prob_forecasts(ts_entry, forecast_entry, asset_name, plot_length=20):
    prediction_intervals = (0.95, 0.99)
    fig, ax = plt.subplots(1, 1, figsize=(10, 7))
    # 为观测值指定标签
    ts_entry[-plot_length:].plot(ax=ax, label="observations")
    # 传入参数让GluonTS生成带标签的预测元素
    forecast_entry.plot(
        intervals=prediction_intervals,
        color='g',
        name="median prediction",
        show_label=True
    )    
    plt.grid(which="both")
    # 直接调用legend()自动收集所有带标签的元素
    plt.legend(loc="upper left")
    plt.title(f'Forecast of {asset_name} series Returns')
    plt.show()

此方法下Matplotlib会自动关联所有带标签的线条和区间,颜色匹配正常,get_legend_handles_labels()也能获取到全部图例项。

方法2:手动构造图例句柄(自定义标签场景)

如果需要自定义标签文本,可以手动收集所有绘图元素的句柄和标签:

def plot_prob_forecasts(ts_entry, forecast_entry, asset_name, plot_length=20):
    prediction_intervals = (0.95, 0.99)
    fig, ax = plt.subplots(1, 1, figsize=(10, 7))
    # 绘制观测值并获取句柄
    obs_line = ts_entry[-plot_length:].plot(ax=ax, label="observations")
    # 绘制预测,开启show_label以生成标签
    forecast_entry.plot(
        intervals=prediction_intervals,
        color='g',
        name="median prediction",
        show_label=True
    )
    # 获取所有句柄和标签
    handles, labels = ax.get_legend_handles_labels()
    # 自定义标签(需与句柄顺序对应)
    custom_labels = ["observations", "median prediction", "99% prediction interval", "95% prediction interval"]
    plt.grid(which="both")
    plt.legend(handles, custom_labels, loc="upper left")
    plt.title(f'Forecast of {asset_name} series Returns')
    plt.show()

注意:fill_between生成的区间句柄会按intervals的顺序添加,自定义标签需对应此顺序。

方法3:临时补丁GluonTS的plot方法

如果不想修改调用代码,可以临时替换GluonTS的plot方法,确保区间元素始终带有标签:

# 在调用绘图代码前执行此补丁(仅当前会话生效)
from gluonts.model.forecast import Forecast
from gluonts.util import maybe

original_plot = Forecast.plot

def patched_plot(self, *, intervals=(0.5, 0.9), ax=None, color=None, name=None, show_label=False):
    import matplotlib.pyplot as plt
    ax = maybe.unwrap_or_else(ax, plt.gca)
    color = maybe.unwrap_or_else(color, lambda: ax._get_lines.get_next_color())
    
    # 绘制中位数,确保有默认标签
    median_line = ax.plot(
        self.index.to_timestamp(),
        self.quantile(0.5),
        color=color,
        label=name if name else "median prediction"
    )[0]
    
    # 为每个区间生成默认标签
    for interval in intervals:
        low = (1 - interval) / 2
        # 强制生成标签,不受show_label限制
        label = f"{int(interval*100)}% prediction interval" 
        ax.fill_between(
            self.index.to_timestamp(),
            self.quantile(low),
            self.quantile(1 - low),
            alpha=0.5 - interval / 3,
            facecolor=color,
            label=label
        )
    return ax

Forecast.plot = patched_plot

之后调用原绘图代码即可正常生成带正确颜色和标签的图例。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 05:23:16