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

MNIST手写数字识别中验证损失优于训练损失的原因探究

MNIST手写数字识别训练与验证损失异常问题

我们用Keras实现了一个简单全连接神经网络,用于MNIST手写数字识别(训练集包含60000张28×28像素图像)。将训练集的10%随机划分为验证集后,得到的训练与验证损失曲线显示前几轮验证损失优于训练损失,这与预期不符。

我们已经排除了以下常见原因:

  • Dropout:未使用任何Dropout(详见下方代码)
  • 训练损失计算时机:无法解释第2轮训练损失(0.29)差于第1轮验证损失(0.26)
  • 验证集选择偏差:尝试不同验证集后仍出现相同现象

实现代码

import numpy as np
from matplotlib import pyplot as plt
import keras

RANDOM_SEED = 0xdeadbeef

from keras.datasets import mnist

# download official train and test sets
(x_train, y_train), (x_test, y_test) = mnist.load_data()

# the original image data are 8bit integers, we normalize them to floats in [0., 1.]
x_train = x_train.astype('float32') / 255
x_test = x_test.astype('float32') / 255


y_train = keras.utils.to_categorical(y_train)
y_test = keras.utils.to_categorical(y_test)

print("train samples: ", x_train.shape[0])
print("test samples: ", x_test.shape[0])

print("shape of one sample: ", x_train[0].shape)



# define a simple feed-forward neural network.

from keras.models import Sequential
from keras.layers import Input, Flatten, Dense

model = Sequential()
model.add(Input(shape=(28, 28)))  # define input shape, here 28x28 images
model.add(Flatten())              # flatten 28x28 images to 784-dimensional vectors
model.add(Dense(128, activation="relu"))    # hidden layer with 128 nodes and relu activation
model.add(Dense(10, activation="softmax"))  # output layer with 10 nodes (for 10 classes) and softmax activation

model.summary()

model.compile(loss="categorical_crossentropy", optimizer="sgd", metrics=["accuracy"])

history = model.fit(
  x_train,
  y_train,
  batch_size=16,
  epochs=40,
  validation_split=.1,
)

# helper function to plot the training and validation losses.

def plot_history(history: keras.callbacks.History):
 
  n = len(history.history['loss'])
  plt.plot(np.arange(n), history.history['loss'], label="training loss")
  plt.plot(np.arange(n), history.history['val_loss'], label="validation loss")
  plt.xticks(range(0, n + 1, 2))
  plt.legend()
  plt.show()


plot_history(history)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 03:25:28