MNIST数据集随机采样绘图触发ValueError错误,求兼容修复方案
问题根因
- scikit-learn 0.24及更高版本中,
fetch_openml接口默认返回 pandas 格式的数据集:特征集X为DataFrame类型,标签y为Series类型,并非旧版本默认的numpy数组。你代码中的索引逻辑是针对numpy数组编写的,不符合pandas的索引规则,因此触发报错。 np.argwhere(y == str(i))[0]返回的是shape为(1,)的嵌套数组,直接用于索引pandas DataFrame时,会触发数据长度和索引长度不匹配的校验报错。
修复方案
最简单的兼容方案是直接将数据集转为numpy数组后再执行后续逻辑,只需修改2行代码即可,原有逻辑无需调整:
import sklearn import pandas as pd import matplotlib.pyplot as plt import numpy as np from sklearn.datasets import fetch_openml mnist = fetch_openml('mnist_784') y = mnist.target.values # 转为numpy数组 X = mnist.data.astype('float64').values # 转为numpy数组 fig, ax = plt.subplots(2, 5) ax = ax.flatten() for i in range(10): im_idx = np.argwhere(y == str(i))[0] print(im_idx) plottable_image = np.reshape(X[im_idx], (28, 28)) ax[i].imshow(plottable_image, cmap='gray_r') plt.show()
如果需要保留pandas格式操作,可以调整索引逻辑如下:
for i in range(10): # 直接取第一个匹配标签的索引 im_idx = y[y == str(i)].index[0] plottable_image = np.reshape(X.loc[im_idx].values, (28, 28)) ax[i].imshow(plottable_image, cmap='gray_r')
可选优化(随机采样)
你当前代码每次取的是每个数字的第一个样本,如果要实现随机采样,可将im_idx的取值逻辑修改为:
im_idx = np.random.choice(np.argwhere(y == str(i)).flatten())
内容的提问来源于stack exchange,提问作者antima99
相关产品推荐
相关产品推荐

