卷积神经网络(CNN)入门教程

Chen Xi
Chen Xi

卷积神经网络(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。增强方式必须根据标签语义选择,并且只应用于训练数据。

聚类和降维等不依赖标签的方法,可参考 无监督学习入门

参考资料