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

TensorFlow实现CIFAR-10图像分类器准确率极低问题求助

CIFAR-10分类器准确率极低、交叉熵过高的问题排查与修复

Hey there! Let's dig into why your CIFAR-10 classifier is stuck with ~10% accuracy (basically random guessing for 10 classes) and sky-high cross-entropy. I spotted several key issues in your code that are causing this—let's break them down one by one and fix things up.

1. Critical: Pixel Values Aren't Normalized

CIFAR-10's raw pixel data is stored as integers between 0 and 255, but your network expects scaled inputs (like MNIST's 0-1 range). Feeding 0-255 directly into ReLU layers leads to extremely large activation values, which messes up gradient flow and makes the model unable to learn. This is almost certainly the biggest reason for your high cross-entropy and low accuracy.

Fix: Divide your input data by 255.0 to scale it to the [0, 1] range.

2. You're Reusing the Same Mini-Batches Over and Over

Your training loop uses i%100 to index into data_batch_1, which means you're just cycling through the same 100 samples repeatedly. You're not iterating through the full dataset, and you're ignoring the other 4 training batches entirely—your model has way too little diverse data to learn meaningful patterns.

Fix: Load all CIFAR-10 training batches, shuffle the combined dataset, and iterate through unique mini-batches each epoch.

3. Output Layer Bias Initialization is Too Aggressive

Your final layer bias B5 = tf.Variable(tf.ones([10])) initializes all biases to 1.0, which pushes the initial logits to high values. This makes the softmax output extremely skewed, leading to massive initial cross-entropy that's hard for the optimizer to recover from.

Fix: Initialize B5 the same way as your other biases: tf.ones([10])/10.

4. Learning Rate Might Be Too High (After Normalization)

While Adam handles learning rates well, 0.003 is on the higher side for normalized inputs. Once you fix the data scaling, tuning this down to 0.001 can help stabilize training.


Fixed Code Implementation

Here's the revised code with all these fixes applied:

import tensorflow as tf
import numpy as np
import os

# Disable unnecessary warnings
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2'

def unpickle(file):
    import pickle
    with open(file, 'rb') as fo:
        dict = pickle.load(fo, encoding='bytes')
    return dict

def load_full_cifar_train_data(data_dir):
    # Load all 5 training batches
    all_data = []
    all_labels = []
    for i in range(1, 6):
        batch = unpickle(f"{data_dir}/data_batch_{i}")
        all_data.append(batch[b'data'])
        all_labels.append(batch[b'labels'])
    # Combine into single arrays
    data = np.concatenate(all_data, axis=0)
    labels = np.concatenate(all_labels, axis=0)
    # Normalize pixel values to [0, 1]
    data = data.astype(np.float32) / 255.0
    return data, labels

def formatLabels(labels):
    # Convert labels to one-hot encoding
    return tf.keras.utils.to_categorical(labels, 10)

# Configuration
DATA_DIR = 'D:/cifar-10-python/cifar-10-batches-py'
BATCH_SIZE = 100
EPOCHS = 10

# Load and prepare data
train_data, train_labels = load_full_cifar_train_data(DATA_DIR)
train_labels_onehot = formatLabels(train_labels)
num_samples = train_data.shape[0]

# Build the model
tf.set_random_seed(0)
L = 200
M = 100
N = 60
O = 30

X = tf.placeholder(tf.float32, [None, 3072])
Y_ = tf.placeholder(tf.float32, [None, 10])

W1 = tf.Variable(tf.truncated_normal([3072, L], stddev=0.1))
B1 = tf.Variable(tf.ones([L])/10)
W2 = tf.Variable(tf.truncated_normal([L, M], stddev=0.1))
B2 = tf.Variable(tf.ones([M])/10)
W3 = tf.Variable(tf.truncated_normal([M, N], stddev=0.1))
B3 = tf.Variable(tf.ones([N])/10)
W4 = tf.Variable(tf.truncated_normal([N, O], stddev=0.1))
B4 = tf.Variable(tf.ones([O])/10)
W5 = tf.Variable(tf.truncated_normal([O, 10], stddev=0.1))
B5 = tf.Variable(tf.ones([10])/10)  # Fixed bias initialization

Y1 = tf.nn.relu(tf.matmul(X, W1) + B1)
Y2 = tf.nn.relu(tf.matmul(Y1, W2) + B2)
Y3 = tf.nn.relu(tf.matmul(Y2, W3) + B3)
Y4 = tf.nn.relu(tf.matmul(Y3, W4) + B4)
Ylogits = tf.matmul(Y4, W5) + B5
Y = tf.nn.softmax(Ylogits)

# Loss and optimizer
cross_entropy = tf.nn.softmax_cross_entropy_with_logits(logits=Ylogits, labels=Y_)
cross_entropy = tf.reduce_mean(cross_entropy) * 100
accuracy = tf.reduce_mean(tf.cast(tf.equal(tf.argmax(Y, 1), tf.argmax(Y_, 1)), tf.float32))
train_step = tf.train.AdamOptimizer(0.001).minimize(cross_entropy)  # Tuned learning rate

# Training loop
init = tf.global_variables_initializer()
with tf.Session() as sess:
    sess.run(init)
    for epoch in range(EPOCHS):
        # Shuffle data each epoch to prevent order bias
        permutation = np.random.permutation(num_samples)
        shuffled_data = train_data[permutation]
        shuffled_labels = train_labels_onehot[permutation]
        
        total_acc = 0.0
        total_loss = 0.0
        num_batches = num_samples // BATCH_SIZE
        
        for batch_idx in range(num_batches):
            start = batch_idx * BATCH_SIZE
            end = start + BATCH_SIZE
            batch_data = shuffled_data[start:end]
            batch_labels = shuffled_labels[start:end]
            
            _, acc, loss = sess.run([train_step, accuracy, cross_entropy], 
                                   feed_dict={X: batch_data, Y_: batch_labels})
            
            total_acc += acc
            total_loss += loss
        
        # Print epoch stats
        avg_acc = total_acc / num_batches
        avg_loss = total_loss / num_batches
        print(f"Epoch {epoch+1}/{EPOCHS}")
        print(f"Average Accuracy: {avg_acc:.4f}")
        print(f"Average Cross-Entropy: {avg_loss:.4f}\n")

Additional Notes

  • Add a Test Set: After training, load the test_batch from CIFAR-10 and evaluate your model on it to check generalization.
  • Consider Regularization: CIFAR-10 is more complex than MNIST—adding dropout layers (tf.nn.dropout) or L2 regularization can help prevent overfitting once your model starts learning.
  • Data Augmentation: For better performance, you can add simple augmentations like random flips or shifts to your training data (though this requires extra preprocessing).

With these fixes, you should see accuracy climb well above 10% and cross-entropy drop significantly within the first few epochs.

内容的提问来源于stack exchange,提问作者Peter

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:49:28