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

如何基于布尔掩码正确过滤Cifar10等多维Numpy数组数据

Numpy多维数组布尔掩码过滤Cifar10数据集正确方法

问题原因

Cifar10数据集加载得到的标签数组trainy/testy默认是形状为(样本数, 1)的二维数组,直接通过标签比较生成的布尔掩码trainMask/testMask同样是二维结构,形状为(样本数, 1),并非和样本轴长度匹配的一维布尔数组。用这个二维掩码直接索引四维的图像特征数组时,numpy会按照高维索引规则匹配维度,无法得到按样本筛选的预期结果。
而标签数组本身是二维,用同形状的二维掩码索引时恰好能返回符合预期的结果,才会出现标签筛选正常、特征筛选失效的现象。

正确实现方法

两种方案都可以实现正确过滤,按需选择即可:

方案1:先将掩码压缩为一维数组再索引

用squeeze()方法去掉掩码长度为1的维度,得到和样本数等长的一维布尔数组,之后直接按原写法索引即可:

from keras.datasets import cifar10
# 加载数据集
(trainX, trainy), (testX, testy) = cifar10.load_data()

# 生成筛选掩码
trainMask = (trainy == 1) | (trainy == 8) | (trainy == 9)
testMask  = (testy == 1)  | (testy == 8)  | (testy == 9)

# 将二维掩码压缩为一维
trainMask = trainMask.squeeze()
testMask = testMask.squeeze()

# 执行筛选
trainX_filtered = trainX[trainMask]
trainy_filtered = trainy[trainMask]
testX_filtered = testX[testMask]
testy_filtered = testy[testMask]

方案2:索引时明确指定样本轴应用掩码

不需要修改掩码形状,索引时明确在第0轴(样本轴)传入展平的掩码,其余维度全部选中即可:

from keras.datasets import cifar10
# 加载数据集
(trainX, trainy), (testX, testy) = cifar10.load_data()

# 生成筛选掩码
trainMask = (trainy == 1) | (trainy == 8) | (trainy == 9)
testMask  = (testy == 1)  | (testy == 8)  | (testy == 9)

# 明确指定第0轴用掩码筛选,其余维度全选
trainX_filtered = trainX[trainMask.ravel(), :, :, :]
trainy_filtered = trainy[trainMask.ravel(), :]
testX_filtered = testX[testMask.ravel(), :, :, :]
testy_filtered = testy[testMask.ravel(), :]

结果验证

筛选完成后打印形状即可确认结果正确:

print('过滤后训练集: X=%s, y=%s' % (trainX_filtered.shape, trainy_filtered.shape))
print('过滤后测试集: X=%s, y=%s' % (testX_filtered.shape, testy_filtered.shape))

预期输出:

过滤后训练集: X=(15000, 32, 32, 3), y=(15000, 1)
过滤后测试集: X=(3000, 32, 32, 3), y=(3000, 1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 05:24:28