Matplotlib子图图例样式不匹配问题排查与解决
问题:Matplotlib全局图例样式匹配异常(重名参数引发)
我用Matplotlib绘制子图对比Testdata1和Testdata2中的多组参数,为统一管理图例,采用fig.legend()在图层面添加全局图例,而非给每个子图单独设置。
代码可正常运行,但偶尔出现图例样式错乱:Dataset2本该显示黑色折线+方形标记,却变成红色折线+圆形标记。尝试用axs[row, col].get_legend_handles_labels()提取图例句柄,返回空值。
排查后确认根源:输入的测量文件存在重名参数,当第一个子图绘制的是重名参数时,图例就会异常;若第一个子图使用唯一参数,图例显示正常。目前临时方案是确保第一个子图用唯一参数,但需要更可靠的解决办法。
原始代码
import pandas as pd import numpy as np import matplotlib.pyplot as plt import math paths = ['Testdata1.xlsx', 'Testdata2.xlsx'] Sheets = [['List1'],['List2']] # define the labels of the plot Labels = ['Dataset1','Dataset2'] Stil = ['ro-', 'ks-' ] # define the x-axis parameter, label and range x_axis ="P_EFF_ME" # Normname Samme x-akse alle plot x_axis_label = "x-axes title" # define the parameters to plot ParamTitles = ["A-heading", "B-heading" , "C-heading" , "D-heading" ] # Plot heading ParamLabels = ["A" , " B", " C", " D" ] # Y-axis text ParamSelect = ["A_data", "B_data" , "C_data" , "D_data" ] # Parametername # calculate the number of rows and columns for the subplots num_plots = len(ParamSelect) # number of plots is the same as numbers of elements in ParamSelect num_cols = min(num_plots, 2) # number of columns is defined as the smallest of either num_plots or 2. If num_plots = 3, then num_cols will still be 2. num_rows = math.ceil(num_plots / num_cols) # number of rows is found by rounding up num_plots/num_cols. If num_plots = 5 and num_cols = 2, then num_plots/num_cols = 2,5. math.ceil rounds up to 3. fig, axs = plt.subplots(num_rows, num_cols, figsize=(9, 3 * num_rows))#,layout='constrained') fig.subplots_adjust(hspace=0.5,wspace=0.3,top=0.9,bottom=0.075) # adjusts margins for k, path in enumerate(paths): for j, sheet in enumerate(Sheets[k]): # read in the data for the current sheet and arranges it as desired for further plotting df = pd.read_excel(path, sheet_name=sheet).T df.index=np.arange(0,df.shape[0],1) df.drop([0,2,3], axis=0,inplace=True) df.columns=df.iloc[0,:] df.index=np.arange(0,df.shape[0],1) df.drop([0], axis=0,inplace=True) for i, parameter in enumerate(ParamSelect): # determine the row and column index for the subplot row = i // num_cols col = i % num_cols # plot the data on the appropriate subplot and add a title axs[row, col].plot(df[x_axis], df[parameter],Stil[k]) axs[row, col].set_xlabel(x_axis_label) axs[row, col].set_ylabel(ParamLabels[i]) axs[row, col].set_title(ParamTitles[i]) # add title fig.legend(labels=Labels, loc='outside upper center', bbox_to_anchor=(0.5, 1), ncol=len(Labels), frameon=False) plt.show()
复现问题的简化代码
import pandas as pd import numpy as np import matplotlib.pyplot as plt import math paths = ['Testdata1', 'Testdata2'] Sheets = [['List1'],['List2']] # define the labels of the plot Labels = ['Dataset1','Dataset2'] Stil = ['ro-', 'ks-'] Testdata1=pd.DataFrame({"B0_data":["B_data",3,3,4,1],"C0_data":["C_data",5,3,6,8],"A1_data":["A_data",1,3,5,7],"D0_data":["D_data",9,5,6,4],"":[0,0,0,0,0],"A0_data":["A_data",1.1,3.1,5.1,7.1],"P0_EFF_ME":["P_EFF_ME",0,10,20,30]}) Testdata2=pd.DataFrame({"B0_data":["B_data",2.4,2.4,3.2,0.8],"C0_data":["C_data",4,2.4,4.8,6.4],"A1_data":["A_data",0.8,2.4,4,5.6],"D0_data":["D_data",7.2,4,4.8,3.2],"A0_data":["A_data",0.8,2.4,4,5.6],"P0_EFF_ME":["P_EFF_ME",0,8,16,24]}) # define the x-axis parameter, label and range x_axis ="P_EFF_ME" # Normname Samme x-akse all plots x_axis_label = "x-axes title" # define the parameters to plot| ParamTitles = ["A-heading", "B-heading" , "C-heading" , "D-heading" ] # Plot heading ParamLabels = ["A" , " B", " C", " D" ] # Y-axis text ParamSelect = ["A_data", "B_data" , "C_data" , "D_data" ] # Parametername # If for example B_data is used first, the legends gets correct #ParamSelect = ["B_data", "A_data" , "C_data" , "D_data" ] # Parametername # calculate the number of rows and columns for the subplots num_plots = len(ParamSelect) # number of plots is the same as numbers of elements in ParamSelect num_cols = min(num_plots, 2) # number of columns is defined as the smallest of either num_plots or 2. If num_plots = 3, then num_cols will still be 2. num_rows = math.ceil(num_plots / num_cols) # number of rows is found by rounding up num_plots/num_cols. If num_plots = 5 and num_cols = 2, then num_plots/num_cols = 2,5. math.ceil rounds up to 3. fig, axs = plt.subplots(num_rows, num_cols, figsize=(9, 3 * num_rows))#,layout='constrained') fig.subplots_adjust(hspace=0.5,wspace=0.3,top=0.9,bottom=0.075) # adjusts margins for k, path in enumerate(paths): for j, sheet in enumerate(Sheets[k]): # read in the data df = globals()[path] df.index=np.arange(0,df.shape[0],1) df.columns=df.iloc[0,:] # Set column name df.index=np.arange(0,df.shape[0],1) df.drop([0], axis=0,inplace=True) #Sletter opprinnelig Normname kolonne for i, parameter in enumerate(ParamSelect): # determine the row and column index for the subplot row = i // num_cols col = i % num_cols # plot the data on the appropriate subplot and add a title axs[row, col].plot(df[x_axis], df[parameter],Stil[k]) axs[row, col].set_xlabel(x_axis_label) axs[row, col].set_ylabel(ParamLabels[i]) axs[row, col].set_title(ParamTitles[i]) # add title fig.legend(labels=Labels, loc='outside upper center', bbox_to_anchor=(0.5, 1), ncol=len(Labels), frameon=False) plt.show()
可靠解决方案
方法1:手动创建图例句柄(最直接)
既然两个数据集的样式是固定的,直接手动创建匹配样式的图例句柄,完全不依赖绘图时的自动收集,彻底规避重名参数的影响。
核心修改代码:
from matplotlib.lines import Line2D # 在绘图完成后,手动生成图例的handles handles = [] for style in Stil: # 解析样式字符串的颜色、标记、线型 color = style[0] marker = style[1] linestyle = style[2] if len(style) > 2 else '-' # 创建自定义线条句柄 handle = Line2D([], [], color=color, marker=marker, linestyle=linestyle, markersize=8) handles.append(handle) # 使用手动创建的handles生成全局图例 fig.legend(handles=handles, labels=Labels, loc='outside upper center', bbox_to_anchor=(0.5, 1), ncol=len(Labels), frameon=False)
方法2:显式指定label并收集首个子图的有效句柄
在绘图时给每个数据集的线条显式添加label,仅在第一个子图添加即可(避免重复),然后从第一个子图提取句柄和标签,再生成全局图例。
核心修改代码:
# 修改绘图循环部分 for k, path in enumerate(paths): for j, sheet in enumerate(Sheets[k]): # ...(数据读取部分不变) for i, parameter in enumerate(ParamSelect): row = i // num_cols col = i % num_cols # 仅在第一个子图添加label,其余子图只绘图 if i == 0: axs[row, col].plot(df[x_axis], df[parameter], Stil[k], label=Labels[k]) else: axs[row, col].plot(df[x_axis], df[parameter], Stil[k]) # ...(子图设置部分不变) # 从第一个子图提取有效句柄和标签 handles, labels = axs[0, 0].get_legend_handles_labels() fig.legend(handles=handles, labels=labels, loc='outside upper center', bbox_to_anchor=(0.5, 1), ncol=len(Labels), frameon=False)
方法3:预处理数据合并重名参数
如果重名参数是同一指标的重复测量,可以在数据读取阶段合并这些重名列,避免绘图时重复绘制相同参数导致的样式冲突。例如对重名列取均值或保留其中一列:
# 在数据读取后添加重名列合并逻辑 df = df.T # ...(原有数据处理步骤) # 按列名分组,取均值合并重名列 df = df.groupby(df.columns, axis=1).mean()
修改后的简化代码(方法1实现)
import pandas as pd import numpy as np import matplotlib.pyplot as plt import math from matplotlib.lines import Line2D paths = ['Testdata1', 'Testdata2'] Sheets = [['List1'],['List2']] # define the labels of the plot Labels = ['Dataset1','Dataset2'] Stil = ['ro-', 'ks-'] Testdata1=pd.DataFrame({"B0_data":["B_data",3,3,4,1],"C0_data":["C_data",5,3,6,8],"A1_data":["A_data",1,3,5,7],"D0_data":["D_data",9,5,6,4],"":[0,0,0,0,0],"A0_data":["A_data",1.1,3.1,5.1,7.1],"P0_EFF_ME":["P_EFF_ME",0,10,20,30]}) Testdata2=pd.DataFrame({"B0_data":["B_data",2.4,2.4,3.2,0.8],"C0_data":["C_data",4,2.4,4.8,6.4],"A1_data":["A_data",0.8,2.4,4,5.6],"D0_data":["D_data",7.2,4,4.8,3.2],"A0_data":["A_data",0.8,2.4,4,5.6],"P0_EFF_ME":["P_EFF_ME",0,8,16,24]}) # define the x-axis parameter, label and range x_axis ="P_EFF_ME" # Normname Samme x-akse all plots x_axis_label = "x-axes title" # define the parameters to plot ParamTitles = ["A-heading", "B-heading" , "C-heading" , "D-heading" ] # Plot heading ParamLabels = ["A" , " B", " C", " D" ] # Y-axis text ParamSelect = ["A_data", "B_data" , "C_data" , "D_data" ] # Parametername # calculate the number of rows and columns for the subplots num_plots = len(ParamSelect) num_cols = min(num_plots, 2) num_rows = math.ceil(num_plots / num_cols) fig, axs = plt.subplots(num_rows, num_cols, figsize=(9, 3 * num_rows)) fig.subplots_adjust(hspace=0.5,wspace=0.3,top=0.9,bottom=0.075) # adjusts margins for k, path in enumerate(paths): for j, sheet in enumerate(Sheets[k]): # read in the data df = globals()[path] df.index=np.arange(0,df.shape[0],1) df.columns=df.iloc[0,:] # Set column name df.index=np.arange(0,df.shape[0],1) df.drop([0], axis=0,inplace=True) for i, parameter in enumerate(ParamSelect): row = i // num_cols col = i % num_cols # plot the data on the appropriate subplot axs[row, col].plot(df[x_axis], df[parameter],Stil[k]) axs[row, col].set_xlabel(x_axis_label) axs[row, col].set_ylabel(ParamLabels[i]) axs[row, col].set_title(ParamTitles[i]) # 手动创建图例句柄 handles = [] for style in Stil: color = style[0] marker = style[1] linestyle = style[2] if len(style) > 2 else '-' handle = Line2D([], [], color=color, marker=marker, linestyle=linestyle, markersize=8) handles.append(handle) # 添加全局图例 fig.legend(handles=handles, labels=Labels, loc='outside upper center', bbox_to_anchor=(0.5, 1), ncol=len(Labels), frameon=False) plt.show()
内容的提问来源于stack exchange,提问作者KVa
相关产品推荐
相关产品推荐

