如何修正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无法解决图例问题
问题分析
- 图例颜色异常:手动指定
legend列表时,Matplotlib无法将自定义标签与fill_between生成的区间图形关联,导致颜色匹配混乱。 - 图例项无法获取:默认
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
相关产品推荐
相关产品推荐

