Skip to content

Keras - Applications

Keras - 预训练模型(Applications 模块)

Section titled “Keras - 预训练模型(Applications 模块)”

tensorflow.keras.applications 模块提供了方便访问带有预训练权重的深度学习模型。这些模型通常在 ImageNet 等大型数据集上训练,对于各种任务都非常有用:

  • 预测: 直接使用模型对新图像进行分类。
  • 特征提取: 使用模型的卷积基(convolutional base,不包含最终的分类层)从图像中提取有意义的特征,这些特征随后可作为另一个模型的输入。
  • 迁移学习 / 微调: 将预训练模型适应于使用较小数据集的新相关任务。这涉及解冻后面的一些层,并在您的特定数据上重新训练它们,从而利用原始训练中学到的特征。

一个训练好的模型由两部分组成:模型架构(层及其连接)和模型权重(训练过程中学习到的参数)。权重文件通常很大,在您首次使用 weights='imagenet' 实例化模型时会自动下载。

tf.keras.applications 中一些流行的预训练模型包括:

  • VGG16, VGG19
  • ResNet50, ResNet101, ResNet152 (及 ResNetRS 变体)
  • InceptionV3
  • InceptionResNetV2
  • MobileNet, MobileNetV2, MobileNetV3Small, MobileNetV3Large
  • DenseNet121, DenseNet169, DenseNet201
  • NASNetLarge, NASNetMobile
  • EfficientNetB0-B7, EfficientNetV2B0-B3, EfficientNetV2S, M, L
  • Xception

使用 tensorflow.keras.applications 中相应的函数加载这些模型非常简单。weights='imagenet' 参数会自动获取从 ImageNet 数据集学习到的权重。

import tensorflow as tf
# Example: Load VGG16 with ImageNet weights
# The weights file will be downloaded on first use
vgg_model = tf.keras.applications.VGG16(weights='imagenet')
# Example: Load MobileNetV2 with ImageNet weights
mobilenet_model = tf.keras.applications.MobileNetV2(weights='imagenet')
# Example: Load ResNet50
resnet_model = tf.keras.applications.ResNet50(weights='imagenet')
# Example: Load InceptionV3
inception_model = tf.keras.applications.InceptionV3(weights='imagenet')

加载模型时的常用参数:

  • weights:权重的来源(‘imagenet’ 或 None,或权重文件的路径)。None 会随机初始化权重。
  • include_top:是否包含最终的全连接分类层。对于特征提取或迁移学习(您将添加自己的分类器),请将其设置为 False。
  • input_shape:如果 include_top 为 False,可选的形状元组(例如,(224, 224, 3))。如果未指定,模型可以处理任意大小的输入。

一旦加载,这些模型可以直接用于预测等任务,或作为更复杂应用的基础,这将在后面的章节中探讨。