如何unpickle文件以将直接下载的CIFAR-10数据集加载至CNN?
解决CIFAR-10手动下载后的数据加载问题(Unpickle操作)
我完全懂你遇到的麻烦——Keras自动下载CIFAR-10经常因为网络问题卡壳,手动下载后又被「unpickle」这个陌生术语难住。别慌,我一步步帮你把这些本地batch文件转换成CNN能直接用的数据格式。
先简单说下unpickle:CIFAR-10的数据集是用Python的pickle模块序列化保存的二进制文件,unpickle就是把这些序列化文件还原成Python可直接操作的字典、数组等结构的过程。
第一步:编写完整的Unpickle函数
你找到的代码片段不完整,这里给你适配Python3的完整版(因为CIFAR-10的pickle文件是Python2格式保存的,Python3必须指定编码才能避免报错):
import pickle import numpy as np def unpickle(file): # 以二进制读模式打开文件,指定encoding适配Python2格式 with open(file, 'rb') as fo: data_dict = pickle.load(fo, encoding='bytes') return data_dict
第二步:加载训练集和测试集
假设你把所有batch文件(data_batch_1到data_batch_5、test_batch、batches.meta)都放在了./cifar-10-batches-py/目录下,执行以下代码整合数据:
# 初始化训练集容器 train_data = [] train_labels = [] # 循环加载5个训练batch for i in range(1, 6): batch = unpickle(f'./cifar-10-batches-py/data_batch_{i}') # 追加每个batch的图片数据 train_data.append(batch[b'data']) # 追加每个batch的标签数据 train_labels.append(batch[b'labels']) # 合并成numpy数组,方便后续处理 train_data = np.concatenate(train_data) train_labels = np.concatenate(train_labels) # 加载测试集 test_batch = unpickle('./cifar-10-batches-py/test_batch') test_data = test_batch[b'data'] test_labels = np.array(test_batch[b'labels']) # 加载类别名称(可选,用来查看图片对应的类别) meta = unpickle('./cifar-10-batches-py/batches.meta') class_names = [name.decode('utf-8') for name in meta[b'label_names']]
第三步:预处理成CNN需要的格式
CIFAR-10原始数据是(样本数, 3072)的一维数组(32323=3072),我们需要把它转成CNN常用的(样本数, 32, 32, 3)三维图片格式,同时做归一化处理:
# 调整数据形状:从(样本数, 3072)转成(样本数, 3, 32, 32),再转成TensorFlow常用的(样本数, 32, 32, 3) train_data = train_data.reshape(-1, 3, 32, 32).transpose(0, 2, 3, 1) test_data = test_data.reshape(-1, 3, 32, 32).transpose(0, 2, 3, 1) # 归一化:把像素值从0-255缩放到0-1区间,提升模型训练效率 train_data = train_data.astype('float32') / 255.0 test_data = test_data.astype('float32') / 255.0 # 标签转成one-hot编码(如果你的CNN使用categorical_crossentropy损失函数的话) from keras.utils import to_categorical train_labels = to_categorical(train_labels, 10) test_labels = to_categorical(test_labels, 10)
第四步:验证数据是否加载正确
可以随便取一张图片可视化,确认数据没问题:
import matplotlib.pyplot as plt # 取训练集第一张图展示 plt.imshow(train_data[0]) plt.title(class_names[np.argmax(train_labels[0])]) plt.show()
这样处理后,train_data、train_labels、test_data、test_labels就可以直接喂给你的CNN模型开始训练了!
内容的提问来源于stack exchange,提问作者albert1905
相关产品推荐
相关产品推荐

