CIFAR100数据集子类采样代码问题:按单元素索引提取样本
CIFAR100指定子类子采样代码修正
问题根源
你当前代码的问题在于:遍历y_full时,每次用np.where(y_full==i)会找出所有匹配当前标签i的索引,导致X_full[np.where(y_full==i)]取出该子类的全部样本并一次性添加到列表中,最终列表元素是批量样本数组,而非单个样本。
解决方案1:逐个遍历样本索引(直观易懂)
通过enumerate同时获取每个样本的索引和标签,逐个判断并提取符合条件的样本:
import numpy as np from tensorflow import keras from sklearn.model_selection import train_test_split cifar100 = keras.datasets.cifar100 (X_full, y_full), (X_test_full, y_test_full) = cifar100.load_data(label_mode="fine") classes = [0,1,2,3,4,5,6,8,9,12,15,22,23,26,27,34,36,41,47,54] X_tr_full = [] y_tr_full = [] X_test = [] y_test = [] # 处理训练集 for idx, label in enumerate(y_full): label_val = label[0] # 提取二维标签数组中的标量值 if label_val in classes: X_tr_full.append(X_full[idx]) y_tr_full.append(label_val) # 处理测试集 for idx, label in enumerate(y_test_full): label_val = label[0] if label_val in classes: X_test.append(X_test_full[idx]) y_test.append(label_val) # 可选:转换为numpy数组,适配后续模型训练流程 X_tr_full = np.array(X_tr_full) y_tr_full = np.array(y_tr_full) X_test = np.array(X_test) y_test = np.array(y_test)
解决方案2:Numpy布尔掩码(高效简洁)
利用Numpy向量化操作生成布尔掩码,一次性提取所有符合条件的样本,效率远高于循环遍历:
import numpy as np from tensorflow import keras from sklearn.model_selection import train_test_split cifar100 = keras.datasets.cifar100 (X_full, y_full), (X_test_full, y_test_full) = cifar100.load_data(label_mode="fine") classes = [0,1,2,3,4,5,6,8,9,12,15,22,23,26,27,34,36,41,47,54] # 将二维标签数组展平为一维,简化判断逻辑 y_full_flat = y_full.flatten() y_test_full_flat = y_test_full.flatten() # 生成布尔掩码:标记出标签在指定列表中的样本位置 train_mask = np.isin(y_full_flat, classes) test_mask = np.isin(y_test_full_flat, classes) # 通过掩码直接提取目标样本 X_tr_full = X_full[train_mask] y_tr_full = y_full_flat[train_mask] X_test = X_test_full[test_mask] y_test = y_test_full_flat[test_mask]
内容的提问来源于stack exchange,提问作者Mel Campbell
相关产品推荐
相关产品推荐

