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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 04:54:02