如何用for循环绘制4×5布局的seaborn直方图(无需转DataFrame)
问题描述
希望通过for循环绘制4行5列布局的Seaborn直方图,但当前代码生成独立子图而非合并成一张图。使用make_classification生成数据集,尝试用Matplotlib的subplots和subplot设置布局,但调用sns.displot后子图独立显示。能正常绘制非Seaborn直方图,询问是否无需将数据转换为Pandas DataFrame即可实现合并布局。
原代码:
from sklearn.datasets import make_classification import seaborn as sns import numpy as np import pandas as pd from matplotlib import pyplot as plt X_train,y_train = make_classification(n_samples=500, n_features=20, n_informative=9, n_redundant=0, n_repeated=0, n_classes=10, n_clusters_per_class=1, class_sep=9, flip_y=0.2, #weights=[0.5,0.5], random_state=17) sns.set_style('darkgrid') coeff_to_analyze = np.arange(0,20,1) rows = 4 cols = 5 N_BINS = 60 fig, axes = plt.subplots(rows, cols, figsize=(45,12)) for i in coeff_to_analyze: ax = plt.subplot(rows, cols, i+1) sns.displot(X_train[i, :], bins=60, kde=True) ax.set_title(f'Coefficient {i}') fig.tight_layout() plt.savefig(f'Histogram_test.pdf', bbox_inches='tight') plt.show()
问题原因
sns.displot是Figure级别的绘图函数,每次调用都会自动创建新的Figure对象,不会复用提前创建的子图axes,这是子图独立显示的核心原因。- 代码同时使用
plt.subplots创建axes数组,又用plt.subplot重新获取子图,属于重复操作,逻辑混乱。 - 原代码中
X_train[i, :]取的是第i行数据,实际应该取第i列(对应第i个特征),属于逻辑错误。
解决方案(无需转换为DataFrame)
改用Seaborn的Axes级别函数sns.histplot,直接指定要绘制的目标子图ax参数,复用提前创建的布局即可。修正后的代码如下:
from sklearn.datasets import make_classification import seaborn as sns import numpy as np from matplotlib import pyplot as plt X_train,y_train = make_classification(n_samples=500, n_features=20, n_informative=9, n_redundant=0, n_repeated=0, n_classes=10, n_clusters_per_class=1, class_sep=9, flip_y=0.2, #weights=[0.5,0.5], random_state=17) sns.set_style('darkgrid') coeff_to_analyze = np.arange(0,20,1) rows = 4 cols = 5 N_BINS = 60 # 创建4行5列的子图布局,获取axes数组 fig, axes = plt.subplots(rows, cols, figsize=(45,12)) # 遍历每个特征和对应的子图 for idx, i in enumerate(coeff_to_analyze): # 计算当前子图在布局中的行和列索引 row = idx // cols col = idx % cols ax = axes[row, col] # 使用sns.histplot,指定ax参数绘制到目标子图 sns.histplot(X_train[:, i], bins=60, kde=True, ax=ax) ax.set_title(f'Coefficient {i}') # 可选:调整x轴标签字号,避免重叠 ax.tick_params(axis='x', labelsize=8) fig.tight_layout() plt.savefig(f'Histogram_test.pdf', bbox_inches='tight') plt.show()
关键说明
sns.histplot是Axes级函数,支持通过ax参数指定绘制的目标子图,不会创建新的Figure。- 遍历过程中通过索引计算子图的行、列位置,直接从
axes数组中获取对应的子图对象。 - 修正了原代码中特征索引的错误,改为
X_train[:, i]获取第i个特征的所有样本数据。
内容的提问来源于stack exchange,提问作者Joe
相关产品推荐
相关产品推荐

