卷积神经网络(CNN)擅长处理图像等具有空间结构的数据。本文用 MNIST 手写数字演示一个小型分类模型,重点是完整且可复现的训练流程。
示例于 2026-07-29 按 TensorFlow/Keras 当前接口复核。安装方式和 Python 版本要求请以 TensorFlow 官方页面为准。
1. CNN 的核心组成
- 卷积层:用可学习卷积核提取边缘、纹理和更高层特征。
- 激活函数:为网络加入非线性,常用 ReLU。
- 池化层:压缩空间尺寸,降低计算量。
- 全连接层:把提取到的特征映射为分类结果。
输入尺寸经过卷积与池化后会变化。设计网络时应关注张量形状、参数量和过拟合风险,而不只是增加层数。
2. 准备数据
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15
| import numpy as np import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers
tf.keras.utils.set_random_seed(42)
(x_train, y_train), (x_test, y_test) = keras.datasets.mnist.load_data() x_train = x_train.astype("float32") / 255.0 x_test = x_test.astype("float32") / 255.0
x_train = np.expand_dims(x_train, -1) x_test = np.expand_dims(x_test, -1)
print(x_train.shape, y_train.shape)
|
测试集只用于最终评估,不应反复查看测试结果来调参。
3. 构建并训练模型
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35
| model = keras.Sequential( [ keras.Input(shape=(28, 28, 1)), layers.Conv2D(32, kernel_size=3, activation="relu"), layers.MaxPooling2D(pool_size=2), layers.Conv2D(64, kernel_size=3, activation="relu"), layers.MaxPooling2D(pool_size=2), layers.Flatten(), layers.Dropout(0.3), layers.Dense(10, activation="softmax"), ] )
model.compile( optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"], )
callbacks = [ keras.callbacks.EarlyStopping( monitor="val_loss", patience=2, restore_best_weights=True, ) ]
history = model.fit( x_train, y_train, batch_size=128, epochs=15, validation_split=0.1, callbacks=callbacks, )
|
validation_split 留出训练数据的一部分选择训练轮次;EarlyStopping 在验证损失不再改善时停止,并恢复最佳权重。
4. 评估与预测
1 2 3 4 5 6
| test_loss, test_accuracy = model.evaluate(x_test, y_test, verbose=0) print({"loss": test_loss, "accuracy": test_accuracy})
probabilities = model.predict(x_test[:5], verbose=0) predicted_classes = probabilities.argmax(axis=1) print(predicted_classes)
|
准确率不能反映所有错误模式。实际项目还应查看混淆矩阵、分类别指标和错误样本,并检查训练数据是否与真实使用场景一致。
5. 保存与重新加载
Keras 推荐使用原生 .keras 格式保存完整模型:
1 2
| model.save("mnist_cnn.keras") restored_model = keras.models.load_model("mnist_cnn.keras")
|
部署模型时还应保存输入归一化、类别映射、依赖版本和训练数据说明,避免“模型文件能加载,但输入语义已变化”。
6. 数据增强
对于自然图像,可在模型开头加入与任务语义一致的数据增强层:
1 2 3 4 5 6 7
| augmentation = keras.Sequential( [ layers.RandomFlip("horizontal"), layers.RandomRotation(0.05), layers.RandomZoom(0.1), ] )
|
手写数字不一定适合水平翻转,所以不要把这段配置直接用于 MNIST。增强方式必须根据标签语义选择,并且只应用于训练数据。
聚类和降维等不依赖标签的方法,可参考 无监督学习入门。
参考资料