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

Matplotlib子图调用fig.legend()后图例重复显示问题求助

解决Matplotlib子图网格直方图图例重复问题

编写Matplotlib代码在子图网格中绘制多个直方图时,调用fig.legend()后每个绘图的图例重复显示两次,以下是原代码:

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
sns.set_style('darkgrid')
def get_cmap(n, name='hsv'):
    return plt.cm.get_cmap(name, n)
def isSqrt(n):
   sq_root = int(np.sqrt(n))
   return (sq_root*sq_root) == n
df = pd.read_csv('mpg.csv')
df2 = pd.read_csv('dm_office_sales.csv')
df['miles'] = df2['salary']
numericClassifier = ['int16', 'int32', 'int64', 'float16', 'float32', 'float64']
newdf = df.select_dtypes(numericClassifier)
columns = newdf.columns.tolist()
n = len(columns)
cmap = get_cmap(n)
if(isSqrt(n)):
    nrows = ncols = int(np.sqrt(n))
else:
    ncols = int(np.sqrt(n))
    for i in range(ncols,50):
        if ncols*i >= n:
            nrows = i
            break
        else:
            pass
fig,ax = plt.subplots(nrows,ncols)
count = 0
print(nrows,ncols)
for i in range(0,nrows,1):
    for j in range(0,ncols,1):
        print('ncols = {}'.format(j),'nrows = {}'.format(i),'count = {}'.format(count))
        if count<=n-1:
            plt_new = sns.histplot(df[columns[count]],ax=ax[i,j],facecolor=cmap(count),kde=True,edgecolor='black',label=df[columns[count]].name)
            patches = plt_new.get_children()
            for patch in patches:
                patch.set_alpha(0.8)
            color = patches[0].get_facecolor()
            ax[i,j].set_xlabel('{}'.format(df[columns[count]].name))
            ax[i,j].xaxis.label.set_fontsize(10)
            ax[i,j].xaxis.label.set_fontname('ariel')
            ax[i,j].set(xlabel=None)
            ax[i,j].tick_params(axis='y', labelsize=8)
            count+=1
        else:
            break
    
for i in range(0,nrows,1):
    for j in range(0,ncols,1):
        if not ax[i,j].has_data():
            fig.delaxes(ax[i,j])
        else:
            pass

plt.suptitle('Histograms').set_fontname('ariel')
plt.tight_layout()
fig.legend(loc='upper right')
plt.show()

问题原因

调用sns.histplot()时指定了kde=True,这会同时生成直方图和KDE拟合曲线两个图形元素,且两者都会继承你设置的label参数。fig.legend()会自动收集所有子图中的所有图例元素,最终导致每个字段的图例重复显示两次。

解决方法

有两种简单有效的修复方式:

方式一:移除KDE曲线的图例标签

在绘制每个直方图后,手动将KDE曲线的标签设置为_nolegend_,让Matplotlib忽略该元素的图例显示:

# 在绘制histplot的代码后添加这一行
plt_new.lines[0].set_label("_nolegend_")

方式二:手动收集单一图例元素并创建图例

遍历所有有数据的子图,只收集每个子图对应直方图的图例元素,再统一传入fig.legend():

# 替换原代码中的fig.legend(loc='upper right')
handles = []
labels = []
for ax in fig.axes:
    h, l = ax.get_legend_handles_labels()
    if h:
        # 只取第一个元素(对应直方图)
        handles.append(h[0])
        labels.append(l[0])
fig.legend(handles=handles, labels=labels, loc='upper right')

修改后的完整代码(方式一示例)

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
sns.set_style('darkgrid')
def get_cmap(n, name='hsv'):
    return plt.cm.get_cmap(name, n)
def isSqrt(n):
   sq_root = int(np.sqrt(n))
   return (sq_root*sq_root) == n
df = pd.read_csv('mpg.csv')
df2 = pd.read_csv('dm_office_sales.csv')
df['miles'] = df2['salary']
numericClassifier = ['int16', 'int32', 'int64', 'float16', 'float32', 'float64']
newdf = df.select_dtypes(numericClassifier)
columns = newdf.columns.tolist()
n = len(columns)
cmap = get_cmap(n)
if(isSqrt(n)):
    nrows = ncols = int(np.sqrt(n))
else:
    ncols = int(np.sqrt(n))
    for i in range(ncols,50):
        if ncols*i >= n:
            nrows = i
            break
        else:
            pass
fig,ax = plt.subplots(nrows,ncols)
count = 0
print(nrows,ncols)
for i in range(0,nrows,1):
    for j in range(0,ncols,1):
        print('ncols = {}'.format(j),'nrows = {}'.format(i),'count = {}'.format(count))
        if count<=n-1:
            plt_new = sns.histplot(df[columns[count]],ax=ax[i,j],facecolor=cmap(count),kde=True,edgecolor='black',label=df[columns[count]].name)
            # 移除KDE曲线的图例标签,避免重复
            plt_new.lines[0].set_label("_nolegend_")
            patches = plt_new.get_children()
            for patch in patches:
                patch.set_alpha(0.8)
            color = patches[0].get_facecolor()
            ax[i,j].set_xlabel('{}'.format(df[columns[count]].name))
            ax[i,j].xaxis.label.set_fontsize(10)
            ax[i,j].xaxis.label.set_fontname('arial')  # 修正拼写错误:ariel → arial
            ax[i,j].set(xlabel=None)
            ax[i,j].tick_params(axis='y', labelsize=8)
            count+=1
        else:
            break
    
for i in range(0,nrows,1):
    for j in range(0,ncols,1):
        if not ax[i,j].has_data():
            fig.delaxes(ax[i,j])
        else:
            pass

plt.suptitle('Histograms').set_fontname('arial')  # 修正拼写错误:ariel → arial
plt.tight_layout()
fig.legend(loc='upper right')
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 16:00:51