《Keras 快速上手与 Sequential 模型》
Keras 提供两种建模方式:Sequential 顺序模型(层叠层、快速原型)和 函数式 API(复杂拓扑,如多输入/多输出)。本文以 Sequential 为主介绍安装、建模、训练与评估全流程。
1 安装
# 使用 TensorFlow 自带的 tf.keras(推荐,Keras 2.x 已并入 TensorFlow)
pip install tensorflow
# 或独立安装 Keras
pip install keras
验证:
import tensorflow as tf
print(tf.__version__)
独立
keras老版本(2.2.x)底层可选用 TensorFlow / Theano / CNTK 后端;新环境统一使用tf.keras。
2 Sequential 建模示例
以手写数字识别(MNIST)为例:
from tensorflow import keras
from tensorflow.keras import layers
# 1. 加载数据
(x_train, y_train), (x_test, y_test) = keras.datasets.mnist.load_data()
# 2. 归一化:像素 0-255 -> 0-1,并展平为 784 维
x_train = x_train.reshape(-1, 784).astype("float32") / 255
x_test = x_test.reshape(-1, 784).astype("float32") / 255
# 3. 组装 Sequential 模型
model = keras.Sequential([
layers.Dense(128, activation="relu"),
layers.Dropout(0.2),
layers.Dense(10, activation="softmax"), # 10 分类
])
# 4. 编译:指定优化器、损失、评估指标
model.compile(
optimizer="adam",
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)
# 5. 训练
history = model.fit(x_train, y_train, epochs=5, batch_size=32,
validation_split=0.1)
# 6. 评估与预测
test_loss, test_acc = model.evaluate(x_test, y_test, verbose=0)
print("test accuracy:", test_acc)
pred = model.predict(x_test[:5])
print(pred.argmax(axis=1))
3 常用层与参数
| 层 | 说明 | 常用参数 |
|---|---|---|
Dense |
全连接层 | units、activation |
Conv2D |
二维卷积 | filters、kernel_size、padding、strides |
MaxPooling2D |
池化 | pool_size |
Flatten |
展平,接全连接前必用 | - |
Dropout |
随机丢弃防过拟合 | rate |
Embedding |
词嵌入 | input_dim、output_dim |
LSTM / SimpleRNN |
循环层 | units、return_sequences |
BatchNormalization |
批归一化 | - |
4 模型配置与保存
# 打印结构
model.summary()
# 查看各层参数
layer = model.layers[0]
print(layer.get_weights())
# 保存模型
model.save("mnist.h5") # HDF5
# 或目录格式(含结构 + 权重 + 优化器状态)
model.save("mnist_model")
# 加载模型
from tensorflow import keras
model = keras.models.load_model("mnist.h5")
5 常见问题
input_shape不写 batch dim:第一层用input_shape=(784,),而不是(None, 784)。- 分类任务损失选择:多分类整数标签用
sparse_categorical_crossentropy,独热编码标签用categorical_crossentropy;二分类用binary_crossentropy。 - 显存不足 / 训练太慢:调小
batch_size,或对图像先用Rescaling/ 缩小尺寸。
参考文档
阅读 —
·
全站 —