Skip to content

Keras - Modules

正如前面介绍的,tf.keras 提供了几个核心模块,这些模块包含用于构建、训练和评估神经网络的基本构建块和功能。这些模块提供了预定义的类和函数,封装了常见的深度学习操作。

以下是 tf.keras 中最重要的模块的 breakdown(分解说明):

  • tf.keras.layers: 包含所有层类型(如 Dense、Conv2D、LSTM、Dropout、BatchNormalization 等)。在“Keras - 层”章节中详细介绍。
  • tf.keras.models: 提供了定义模型的方式(Sequential 序列模型,以及用于函数式 API 或模型子类化的 Model 模型)。在“Keras - 模型”章节中介绍。
  • tf.keras.initializers: 用于设置层初始权重的函数(例如 GlorotUniform、HeNormal、Zeros)。有助于模型收敛。
  • tf.keras.regularizers: 用于对层参数或激活值应用正则化惩罚(例如 L1、L2、L1L2)以防止过拟合的类。
  • tf.keras.constraints: 用于在优化期间对层权重施加约束的函数(例如 MaxNorm、NonNeg)。
  • tf.keras.activations: 逐元素应用的激活函数(例如 relu、sigmoid、softmax、tanh)。引入非线性。
  • tf.keras.losses: 用于在训练期间衡量模型误差的损失函数(例如 BinaryCrossentropy 二分类交叉熵、CategoricalCrossentropy 分类交叉熵、MeanSquaredError 均方误差)。
  • tf.keras.metrics: 用于评估模型性能的指标(例如 Accuracy 准确率、Precision 精确率、Recall 召回率、AUC、MeanAbsoluteError 平均绝对误差)。通常在训练和评估期间报告。
  • tf.keras.optimizers: 更新模型权重以最小化损失函数的优化算法(例如 Adam、SGD 随机梯度下降、RMSprop)。
  • tf.keras.callbacks: 可以传递给 model.fit() 的对象,用于在训练过程的不同阶段执行操作(例如,保存模型、提前停止训练、记录到 TensorBoard、调整学习率)。
  • tf.keras.utils: 用于常见任务的实用工具函数,例如数据预处理(to_categorical 转换为 one-hot 编码、load_img 加载图像、img_to_array 图像转数组)、模型可视化(plot_model 绘制模型图)和序列填充(pad_sequences 填充序列)。
  • tf.keras.applications: 预训练的深度学习模型(例如 VGG16、ResNet50、MobileNetV2)。在“Keras - 应用”章节中介绍。
  • tf.keras.backend: 底层抽象函数。虽然现在直接使用它比早期 Keras 时代少见,但如果需要访问后端特定操作,它仍然提供了途径。

让我们深入研究其中一些模块(Optimizers、Losses、Metrics、Callbacks、Utils),因为它们对模型的编译和训练阶段至关重要。

优化器决定了如何根据计算出的损失来更新模型的权重。常见的选择包括:

  • Adam: 一种自适应学习率优化算法,通常是一个不错的默认选择。
  • SGD: 随机梯度下降,常与动量和学习率调度配合使用。
  • RMSprop: 另一种自适应学习率方法。

编译期间的使用:

import tensorflow as tf
optimizer = tf.keras.optimizers.Adam(learning_rate=0.001)
# model.compile(optimizer=optimizer, ...)

您也可以通过字符串标识符指定优化器(例如 'adam'),但实例化类允许进行定制(例如设置学习率)。

损失函数量化了模型的预测与真实目标值之间的差异。选择哪种损失函数取决于具体的任务:

  • BinaryCrossentropy: 用于二分类任务(输出层有 1 个单元并使用 sigmoid 激活)。
  • CategoricalCrossentropy: 用于标签是 one-hot 编码的多分类任务(输出层有 N 个单元并使用 softmax 激活)。
  • SparseCategoricalCrossentropy: 用于标签是整数索引的多分类任务(输出层有 N 个单元并使用 softmax 激活)。
  • MeanSquaredError (mse): 常用于回归任务。
  • MeanAbsoluteError (mae): 回归任务的另一种选择,对离群值不如 MSE 敏感。

编译期间的使用:

loss_fn = tf.keras.losses.CategoricalCrossentropy()
# model.compile(loss=loss_fn, ...)
# Or using string identifiers:
# model.compile(loss='categorical_crossentropy', ...)

指标用于评估模型性能,但不直接影响训练过程(与损失函数不同)。常见的指标包括:

  • Accuracy ('accuracy'): 正确预测的比例(常用于分类)。
  • Precision 精确率, Recall 召回率, AUC: 其他分类指标,提供更细致的评估。
  • MeanAbsoluteError ('mae'): 常在回归中用作指标。
  • RootMeanSquaredError ('RootMeanSquaredError'): 另一种回归指标。

编译期间的使用(通常是一个列表):

metrics_list = [tf.keras.metrics.Accuracy(), tf.keras.metrics.Precision()]
# model.compile(metrics=metrics_list, ...)
# Or using string identifiers:
# model.compile(metrics=['accuracy', 'Precision'], ...)

回调是在模型训练过程的不同点(例如,轮次开始/结束,批次开始/结束)调用的实用工具。它们允许您自动化任务:

  • ModelCheckpoint: 定期保存模型(或仅权重),通常只保存基于某个被监控指标(如验证损失)表现最好的版本。
  • EarlyStopping: 当一个被监控指标(例如验证损失)在一定数量的轮次 (patience) 内没有改善时停止训练。防止过拟合并节省时间。
  • TensorBoard: 记录指标、图可视化等,用于在 TensorBoard UI 中监控训练进度。
  • ReduceLROnPlateau: 当一个指标停止改善时降低学习率。
  • CSVLogger: 将每个轮次的结果记录到一个 CSV 文件。

训练期间的使用:

import tensorflow as tf
# Example Callbacks
early_stopping = tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=10)
model_checkpoint = tf.keras.callbacks.ModelCheckpoint(
filepath='best_model.keras', # Use .keras format
save_best_only=True,
monitor='val_accuracy',
mode='max'
)
tensorboard_callback = tf.keras.callbacks.TensorBoard(log_dir="./logs")
# Pass callbacks to model.fit()
# model.fit(x_train, y_train, epochs=100, validation_data=(x_val, y_val),
# callbacks=[early_stopping, model_checkpoint, tensorboard_callback])

此模块包含用于数据处理和模型检查的有用函数:

  • to_categorical: 将整数类向量转换为 one-hot 编码矩阵。
  • pad_sequences: 填充序列(如文本)到相同的长度。
  • load_img, img_to_array: 用于加载和转换图像文件的实用工具。
  • get_file: 从 URL 下载文件,如果尚未缓存。
  • plot_model: 将 Keras 模型转换为 DOT 格式并保存为文件(需要安装 pydot 和 graphviz)。对于可视化模型架构很有用。

示例用法:

import tensorflow as tf
import numpy as np
# One-hot encode labels
labels = np.array([0, 1, 3, 2])
one_hot_labels = tf.keras.utils.to_categorical(labels, num_classes=4)
# print(one_hot_labels)
# Visualize model (assuming 'model' is a defined Keras model)
# tf.keras.utils.plot_model(model, to_file='model_plot.png', show_shapes=True)

后端模块提供了底层函数。过去,当 Keras 支持多个后端时,这部分用于与后端无关的操作。对于 tf.keras,这些函数通常直接映射到 TensorFlow 操作。现在直接使用较少,因为标准的 TensorFlow 操作 (tf.*) 或更高级别的 Keras 层/函数通常足够。

示例(不太常见):

import tensorflow.keras.backend as K
import tensorflow as tf
# Example: Calculate dot product using backend
tensor1 = tf.constant([[1.0, 2.0], [3.0, 4.0]])
tensor2 = tf.constant([[5.0], [6.0]])
dot_product = K.dot(tensor1, tensor2)
# print(K.eval(dot_product)) # K.eval() evaluates a tensor (needs session in TF1)
# In TensorFlow 2.x, direct evaluation is simpler:
# print(dot_product.numpy())