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

如何在Matplotlib中为多个axvspan设置对应颜色的图例?

解决Matplotlib中axvspan图例颜色不匹配的问题

你遇到的问题根源有两个:一是重复给多个同类型的axvspan设置相同label,Matplotlib生成图例时只会保留最后一个带该label的元素样式;二是手动给plt.legend()传标签的方式,没有和实际绘图元素正确绑定,导致颜色对应错误。

下面提供两种可靠的解决方法:

方案一:仅为同类型的第一个区间设置label

修改axvspan的创建逻辑,只给第一个橙色预测区间和第一个红色崩盘区间设置label,后续同类区间不再设置,然后直接调用plt.legend()自动读取元素的label:

def plot_test_results(df, c, t_start, t_end):
    t_start = [datetime.strptime(t, '%Y-%m-%d') for t in t_start]
    t_end = [datetime.strptime(t, '%Y-%m-%d') for t in t_end]
    for t1, t2 in zip(t_start, t_end):
        gs = gridspec.GridSpec(2, 1, height_ratios=[2.5,1])
        ax=plt.subplot(gs[0])
        y_start = list(df[t1:t2][df.loc[t1:t2, 'y_pred'].diff(-1) < 0].index)
        y_end = list(df[t1:t2][df.loc[t1:t2, 'y_pred'].diff(-1) > 0].index)
        crash_st = list(filter(lambda x: x > t1 and x < t2, c['crash_st']))
        crash_end = list(filter(lambda x: x > t1 and x < t2, c['crash_end']))
        
        # 给价格曲线设置label
        plt.plot(df['price'][t1:t2], color='blue', label='Price') 
        
        # 绘制预测区间:仅第一个设置label
        for idx, (x1, x2) in enumerate(zip(y_start, y_end)):
            label = 'Crash Prediction' if idx == 0 else None
            plt.axvspan(x1, x2, alpha=0.4, color='orange', label=label, zorder=2)
        
        # 绘制崩盘区间:仅第一个设置label
        for idx, (c1, c2) in enumerate(zip(crash_st, crash_end)):
            label = 'Crash' if idx == 0 else None
            plt.axvspan(c1, c2, alpha=0.8, color='red', label=label)
        
        # 自动生成图例
        plt.legend()
        plt.title(test_data + ' ' + model_name +  ', Time period: ' + str(calendar.month_name[t1.month]) + ' ' + str(t1.year) + ' - ' +\
                 str(calendar.month_name[t2.month]) + ' ' + str(t2.year))
        plt.show()

方案二:自定义图例元素(更灵活)

如果同类型区间数量不确定,或者想精准控制图例样式,可以用matplotlib.patches.Patch创建自定义图例项,手动指定颜色和标签:

from matplotlib.patches import Patch

def plot_test_results(df, c, t_start, t_end):
    t_start = [datetime.strptime(t, '%Y-%m-%d') for t in t_start]
    t_end = [datetime.strptime(t, '%Y-%m-%d') for t in t_end]
    for t1, t2 in zip(t_start, t_end):
        gs = gridspec.GridSpec(2, 1, height_ratios=[2.5,1])
        ax=plt.subplot(gs[0])
        y_start = list(df[t1:t2][df.loc[t1:t2, 'y_pred'].diff(-1) < 0].index)
        y_end = list(df[t1:t2][df.loc[t1:t2, 'y_pred'].diff(-1) > 0].index)
        crash_st = list(filter(lambda x: x > t1 and x < t2, c['crash_st']))
        crash_end = list(filter(lambda x: x > t1 and x < t2, c['crash_end']))
        
        plt.plot(df['price'][t1:t2], color='blue') 
        # 绘制区间时无需设置label
        [plt.axvspan(x1, x2, alpha=0.4, color='orange', zorder=2) for x1, x2 in zip(y_start, y_end)]
        [plt.axvspan(c1, c2, alpha=0.8, color='red') for c1, c2 in zip(crash_st, crash_end)]
        
        # 创建自定义图例元素
        legend_elements = [
            Patch(facecolor='blue', label='Price'),
            Patch(facecolor='red', alpha=0.8, label='Crash'),
            Patch(facecolor='orange', alpha=0.4, label='Crash Prediction')
        ]
        # 传入自定义元素生成图例
        plt.legend(handles=legend_elements)
        
        plt.title(test_data + ' ' + model_name +  ', Time period: ' + str(calendar.month_name[t1.month]) + ' ' + str(t1.year) + ' - ' +\
                 str(calendar.month_name[t2.month]) + ' ' + str(t2.year))
        plt.show()

说明

  • 方案一利用Matplotlib自动收集带label的绘图元素生成图例,避免重复label导致的样式覆盖
  • 方案二完全手动控制图例内容,适合复杂场景,确保颜色和标签完全对应

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 17:47:49