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

在TensorFlow中实现CIFAR10时遭遇KeyError问题求助

Fixing KeyError: 'data' When Loading CIFAR10 in Python 3

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 07:35:19