在TensorFlow中实现CIFAR10时遭遇KeyError问题求助
Hey there! The KeyError: 'data' you're hitting is absolutely tied to Python 3's strict distinction between bytes and strings when loading the CIFAR10 pickle files. Let's break down the issue and fix it step by step.
What's Causing the Problem?
The CIFAR10 dataset's pickle files were created in Python 2, where str and bytes types were interchangeable. When you load these files in Python 3 with pickle.load(), all dictionary keys (like 'data' or 'labels') come through as bytes objects (e.g., b'data') instead of regular strings. Trying to access them with string keys like d["data"] naturally throws a KeyError.
Step-by-Step Solutions
1. Update the unpickle Function to Convert Bytes Keys to Strings
Modify your unpickle function to decode all bytes keys in the loaded dictionary to UTF-8 strings. This makes the keys compatible with your string-based access attempts:
def unpickle(file): with open(os.path.join(DATA_PATH, file), 'rb') as fo: dict = pickle.load(fo, encoding='bytes') # Convert all bytes keys to regular strings return {key.decode('utf-8'): value for key, value in dict.items()}
2. Fix a Typo in the next_batch Method
You have a small typo that would cause an error once the KeyError is fixed: sel._i should be self._i. Correct it like this:
def next_batch(self, batch_size): x, y = self.images[self._i:self._i+batch_size], self.labels[self._i:self._i+batch_size] self._i = (self._i + batch_size) % len(self.images) return x, y
3. Add the Missing Matplotlib Import
Your display_cifar function uses plt but doesn't import it. Add this at the top of your code:
import matplotlib.pyplot as plt
Full Corrected Code
Here's the complete code with all fixes applied:
import os import numpy as np import pickle import matplotlib.pyplot as plt # Added missing import class CifarLoader(object): def __init__(self, source_files): self._source = source_files self._i = 0 self.images = None self.labels = None def load(self): data = [unpickle(f) for f in self._source] images = np.vstack([d["data"] for d in data]) n = len(images) self.images = images.reshape(n, 3, 32, 32).transpose(0, 2, 3, 1).astype(float)/255 self.labels = one_hot(np.hstack([d["labels"] for d in data]), 10) return self def next_batch(self, batch_size): x, y = self.images[self._i:self._i+batch_size], self.labels[self._i:self._i+batch_size] self._i = (self._i + batch_size) % len(self.images) # Fixed typo from sel._i to self._i return x, y DATA_PATH = "cifar10" def unpickle(file): with open(os.path.join(DATA_PATH, file), 'rb') as fo: dict = pickle.load(fo, encoding='bytes') # Convert bytes keys to strings return {key.decode('utf-8'): value for key, value in dict.items()} def one_hot(vec, vals=10): n = len(vec) out = np.zeros((n, vals)) out[range(n), vec] = 1 return out class CifarDataManager(object): def __init__(self): self.train = CifarLoader(["data_batch_{}".format(i) for i in range(1, 6)]).load() self.test = CifarLoader(["test_batch"]).load() def display_cifar(images, size): n = len(images) plt.figure() plt.gca().set_axis_off() im = np.vstack([np.hstack([images[np.random.choice(n)] for i in range(size)]) for i in range(size)]) plt.imshow(im) plt.show() d = CifarDataManager() print ("Number of train images: {}".format(len(d.train.images))) print ("Number of train labels: {}".format(len(d.train.labels))) print ("Number of test images: {}".format(len(d.test.images))) print ("Number of test labels: {}".format(len(d.test.labels))) # Fixed duplicate print statement images = d.train.images display_cifar(images, 10)
Why This Works
By decoding bytes keys to strings, we bridge the gap between Python 2's type system (where strings and bytes were interchangeable) and Python 3's strict typing. The other fixes address small oversights that would have caused additional errors once the KeyError was resolved.
内容的提问来源于stack exchange,提问作者Cristi Vlad

