如何在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
相关产品推荐
相关产品推荐

