Python加载CIFAR-10后如何访问字典中的数据?
嘿,刚接触Python遇到这种字典访问的问题太正常了!我来一步步给你拆解怎么从你加载的CIFAR-10数据里取出需要的内容~
首先,你用unpickle加载的每个batch(不管是训练集的data_batch_1到data_batch_5,还是测试集的test_batch)都是一个字节键的字典——因为你用了encoding='bytes'参数,所以字典的键都是类似b'data'这样的字节字符串,这点要注意哦。
第一步:先看看字典里有哪些内容
你可以先打印单个训练batch的键,看看里面都存了什么:
# 取第一个训练batch first_train_batch = train[0] # 打印所有键 print(first_train_batch.keys())
运行后会输出类似这样的结果:
dict_keys([b'batch_label', b'labels', b'data', b'filenames'])
这四个就是CIFAR-10每个batch字典里的固定键,下面逐个讲它们的用途:
第二步:逐个访问字典中的内容
1. 图像数据:b'data'
这个键对应的值是一个形状为(10000, 3072)的numpy数组——每个batch有10000张图,每张图是3072个整数(对应32x32像素的RGB图像,3个通道,每个通道32*32=1024个像素值)。
如果你想把它转换成更直观的图像形状(比如(样本数, 通道数, 高度, 宽度)或者(样本数, 高度, 宽度, 通道数)),可以这样处理:
# 转成(10000, 3, 32, 32)的形状(通道在前,适合PyTorch) train_images = first_train_batch[b'data'].reshape(-1, 3, 32, 32) # 或者转成(10000, 32, 32, 3)的形状(通道在后,适合TensorFlow/OpenCV) train_images = train_images.transpose(0, 2, 3, 1)
2. 类别标签:b'labels'
这个键对应的值是一个长度为10000的列表,每个元素是0-9的整数,分别对应CIFAR-10的10个类别(比如0=飞机,1=汽车,2=鸟类,3=猫,4=鹿,5=狗,6=青蛙,7=马,8=船,9=卡车)。
你可以直接把它转成numpy数组方便后续处理:
train_labels = np.array(first_train_batch[b'labels'])
3. 文件名:b'filenames'
这个是每个图像的原始文件名,是一个字节字符串的列表,比如b'leptodactylus_pentadactylus_s_000004.png',一般做训练的时候用不上,但如果需要对应原始文件可以用它。
4. Batch标识:b'batch_label'
这个就是当前batch的说明,比如第一个训练batch的这个值是b'training batch 1 of 5',测试集的是b'testing batch 1 of 1',主要用来确认加载的是哪个batch。
第三步:合并所有训练batch的数据
因为你加载了5个训练batch,要是想把它们整合成一个完整的训练集,可以这样做:
# 合并所有训练图像数据 all_train_images = np.concatenate([batch[b'data'] for batch in train], axis=0) # 转成合适的形状 all_train_images = all_train_images.reshape(-1, 3, 32, 32).transpose(0, 2, 3, 1) # 合并所有训练标签 all_train_labels = [] for batch in train: all_train_labels.extend(batch[b'labels']) all_train_labels = np.array(all_train_labels)
这样all_train_images就是(50000, 32, 32, 3)的数组,all_train_labels是(50000,)的数组,就是完整的CIFAR-10训练集啦~
测试集的访问方式和单个训练batch完全一样,比如:
test_images = test[b'data'].reshape(-1, 3, 32, 32).transpose(0, 2, 3, 1) test_labels = np.array(test[b'labels'])
要是你还有其他没说出来的困惑,随时告诉我哦!
内容的提问来源于stack exchange,提问作者user9369143

