迁移学习实战:用TensorFlow 2.0快速实现图像分类(附完整代码)

你是否遇到过这样的困境:手头有一个图像分类任务,但标注数据只有寥寥几百张,从头训练一个深度卷积网络不仅耗时耗力,效果还往往不尽如人意?或者,你希望快速验证一个业务想法,却不想在模型训练上花费数周时间?这正是迁移学习大显身手的舞台。它不再是实验室里的理论概念,而是每一位希望高效解决实际问题的开发者工具箱里的必备利器。本文将带你绕过繁琐的理论铺垫,直接进入TensorFlow 2.0的实战世界,通过一个完整的代码案例,手把手教你如何利用预训练模型,在短时间内构建一个强大且精准的图像分类器。无论你是想为电商平台搭建商品识别系统,还是为医疗影像分析提供辅助工具,这里的思路和代码都能为你提供一个坚实的起点。

1. 迁移学习的核心思想:站在巨人的肩膀上

在深入代码之前,我们有必要花几分钟理解迁移学习为何如此有效。想象一下,一个在ImageNet数据集(包含1400万张图片,2万多个类别)上训练过的深度神经网络,它已经学会了识别从“边牧犬”到“汽车轮胎”的无数通用视觉特征,如边缘、纹理、形状和物体部件。这些底层特征,对于绝大多数视觉任务而言,是高度可复用的。

迁移学习的本质,就是复用这些已经习得的、强大的特征提取能力,而非从零开始学习。我们将预训练模型视为一个已经受过多年教育的“专家”,它具备了优秀的“看图说话”基础。我们的任务,是针对自己特定的领域(比如“区分肺炎与正常胸片”),对这个专家进行“微调”或“再培训”,让它快速掌握新领域的专业知识。

这个过程通常分为三步:

  1. 特征提取:将预训练模型(去掉其顶部的分类层)作为一个固定的特征提取器。我们输入自己的图片,模型输出的是高维的特征向量(而非类别概率)。
  2. 分类器训练:在这些提取出的特征之上,我们搭建一个全新的、轻量级的分类器(通常是几个全连接层)。由于特征已经非常具有判别性,这个新分类器可以用很少的数据快速训练好。
  3. 微调:这是更进一步的策略。我们不仅训练新的分类器,还“解冻”预训练模型靠近顶部的若干层,让它们也参与训练,以更好地适应新数据的特点。

为什么这比从头训练好?

  • 数据需求小:你可能只需要几百张,甚至几十张标注图片就能取得不错的效果。
  • 训练速度快:基础特征无需重新学习,训练收敛极快。
  • 性能起点高:基于在大型数据集上验证过的强大模型,效果通常远超随机初始化的模型。

注意:迁移学习并非万能。当你的新任务与预训练模型原始任务(如ImageNet分类)的领域差异极大时(例如,从自然图像到医学断层扫描),效果可能会打折扣。此时,可能需要更激进的微调,或考虑使用在更相关领域预训练的模型。

2. 环境搭建与数据准备

工欲善其事,必先利其器。我们先确保一切就绪。

2.1 安装与导入必要的库

确保你已安装TensorFlow 2.x。我们将使用tensorflow.keras,它是TensorFlow对Keras API的高层实现,极大地简化了模型构建流程。

# 导入核心库
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers, models, applications, preprocessing, optimizers, callbacks
import numpy as np
import matplotlib.pyplot as plt
import os
import pathlib
import warnings
warnings.filterwarnings('ignore')

# 检查TensorFlow版本和GPU是否可用
print("TensorFlow版本:", tf.__version__)
print("GPU是否可用:", tf.config.list_physical_devices('GPU'))

如果你的输出显示有可用的GPU,那么恭喜你,训练速度将得到质的飞跃。TensorFlow会自动利用GPU进行加速。

2.2 组织你的图像数据

一个清晰的数据目录结构是成功的一半。Keras提供了一个极其方便的image_dataset_from_directory函数,能自动从文件夹结构创建标签化的数据集。假设你的数据目录结构如下:

your_dataset/
├── train/
│   ├── class_1/
│   │   ├── img001.jpg
│   │   └── img002.jpg
│   └── class_2/
│       ├── img101.jpg
│       └── img102.jpg
└── validation/
    ├── class_1/
    └── class_2/

这里,class_1class_2是你的两个类别名称。我们使用“猫狗分类”作为贯穿本文的示例,但你可以轻松替换为任何其他类别。

# 设置数据路径
data_dir = pathlib.Path("./your_dataset") # 请修改为你的实际路径
train_dir = data_dir / 'train'
val_dir = data_dir / 'validation'

# 设置训练参数
BATCH_SIZE = 32
IMG_HEIGHT = 224 # 与预训练模型输入尺寸匹配
IMG_WIDTH = 224
AUTOTUNE = tf.data.AUTOTUNE # 用于数据加载的动态调优

# 创建训练数据集
train_ds = preprocessing.image_dataset_from_directory(
    train_dir,
    seed=123,
    image_size=(IMG_HEIGHT, IMG_WIDTH),
    batch_size=BATCH_SIZE,
    label_mode='categorical' # 对于多分类,使用 'categorical' 获得one-hot标签
)

# 创建验证数据集
val_ds = preprocessing.image_dataset_from_directory(
    val_dir,
    seed=123,
    image_size=(IMG_HEIGHT, IMG_WIDTH),
    batch_size=BATCH_SIZE,
    label_mode='categorical'
)

# 获取类别名称
class_names = train_ds.class_names
print(f"发现 {len(class_names)} 个类别: {class_names}")

image_dataset_from_directory 会自动完成图片解码、调整大小和批处理,并生成一个高效的tf.data.Dataset对象。设置label_mode='categorical'是为了适配我们后面使用的分类损失函数。

2.3 数据预处理与增强

数据增强是提升模型泛化能力、防止过拟合的关键技术,尤其是在数据量不大的情况下。我们可以在数据管道中直接插入增强层。

# 定义数据增强管道
data_augmentation = keras.Sequential([
    layers.RandomFlip("horizontal"),
    layers.RandomRotation(0.1),
    layers.RandomZoom(0.1),
    layers.RandomContrast(0.1),
])

# 将增强应用于训练集,并优化数据加载性能
def prepare_dataset(ds, augmentation=False, shuffle=False):
    # 将像素值从[0,255]归一化到[0,1]或[-1,1](取决于预训练模型要求)
    # 这里我们先归一化到[0,1]
    normalization_layer = layers.Rescaling(1./255)
    ds = ds.map(lambda x, y: (normalization_layer(x), y))
    
    if augmentation:
        ds = ds.map(lambda x, y: (data_augmentation(x, training=True), y))
    
    if shuffle:
        ds = ds.shuffle(buffer_size=1000)
    
    # 使用预取重叠数据预处理和模型执行,提升性能
    return ds.prefetch(buffer_size=AUTOTUNE)

train_ds = prepare_dataset(train_ds, augmentation=True, shuffle=True)
val_ds = prepare_dataset(val_ds, augmentation=False, shuffle=False)

现在,我们的数据已经准备就绪,可以高效地流入模型了。

3. 构建迁移学习模型:以EfficientNet为例

TensorFlow Keras提供了丰富的预训练模型,位于tf.keras.applications模块中。我们将选择EfficientNetB0作为基础模型。它在精度和效率之间取得了出色的平衡,是当前非常流行的选择。当然,你也可以尝试ResNet50、MobileNetV2等。

3.1 加载预训练基础模型

关键点在于设置include_top=False,这意味着我们只加载卷积基(特征提取器),而丢弃顶部的全连接分类层。同时,我们要指定输入图片的形状,并加载在ImageNet上预训练的权重。

# 创建基础模型
IMG_SHAPE = (IMG_HEIGHT, IMG_WIDTH, 3)
base_model = applications.EfficientNetB0(
    input_shape=IMG_SHAPE,
    include_top=False, # 不要顶部分类层
    weights='imagenet' # 加载ImageNet预训练权重
)

# 在微调之前,先冻结基础模型的所有层
base_model.trainable = False
print(f"基础模型有 {len(base_model.layers)} 层。")
print(f"前5层是否可训练: {base_model.layers[:5].trainable}")

通过将trainable设置为False,我们冻结了基础模型的所有参数。在接下来的第一步中,我们只训练新添加的顶层分类器。这是一种常见的策略,可以防止预训练权重在初始阶段被破坏。

3.2 添加自定义分类头

现在,我们在基础模型之上构建我们自己的小型分类网络。

# 构建完整模型
inputs = keras.Input(shape=IMG_SHAPE)

# 将数据增强作为模型的一部分(可选,另一种实现方式)
# x = data_augmentation(inputs)

# 基础模型(特征提取器)
# 设置 training=False,确保BatchNorm层在推理模式下运行,即使我们微调时也是如此
x = base_model(inputs, training=False)

# 添加全局平均池化层,将特征图转换为特征向量
# 这比Flatten层参数更少,且能提供一定的平移不变性
x = layers.GlobalAveragePooling2D()(x)

# 添加一个全连接层作为非线性变换
x = layers.Dense(128, activation='relu')(x)
# 添加Dropout层防止过拟合
x = layers.Dropout(0.3)(x)

# 输出层:神经元数量等于类别数,使用softmax激活
outputs = layers.Dense(len(class_names), activation='softmax')(x)

# 构建最终模型
model = keras.Model(inputs, outputs)

# 查看模型架构摘要
model.summary()

GlobalAveragePooling2D层是一个巧妙的设计。它将每个特征图(Channel)的所有值取平均,得到一个单一数值。对于一个有1280个通道的EfficientNetB0顶层,它会输出一个1280维的向量。这极大地减少了参数数量,并整合了空间信息。

3.3 编译模型

在训练之前,我们需要配置学习过程。

# 编译模型
# 使用较低的学习率,因为特征已经很好,我们只需要微调
initial_learning_rate = 0.0001

model.compile(
    optimizer=optimizers.Adam(learning_rate=initial_learning_rate),
    loss='categorical_crossentropy', # 多分类交叉熵损失
    metrics=['accuracy'] # 监控准确率
)

由于基础模型被冻结,我们只训练顶部的几层,因此使用一个较小的学习率(如1e-4)是合适的,以避免在训练初期就迈出太大的步伐,破坏了好的特征。

4. 第一阶段:训练顶层分类器

现在,让我们开始第一轮训练,只训练我们刚刚添加的顶层。

# 设置训练轮次
initial_epochs = 10

# 定义回调函数
# ModelCheckpoint: 保存最佳模型
checkpoint_cb = callbacks.ModelCheckpoint(
    "best_model_phase1.keras",
    save_best_only=True,
    monitor='val_accuracy'
)
# EarlyStopping: 当验证指标不再提升时提前停止训练
early_stopping_cb = callbacks.EarlyStopping(
    monitor='val_accuracy',
    patience=5,
    restore_best_weights=True
)
# ReduceLROnPlateau: 当指标停滞时降低学习率
reduce_lr_cb = callbacks.ReduceLROnPlateau(
    monitor='val_loss',
    factor=0.5,
    patience=3,
    min_lr=1e-7
)

# 开始训练
print("开始第一阶段训练(仅训练顶层)...")
history = model.fit(
    train_ds,
    epochs=initial_epochs,
    validation_data=val_ds,
    callbacks=[checkpoint_cb, early_stopping_cb, reduce_lr_cb]
)

这个阶段通常收敛得很快,可能只需要几轮就能在验证集上达到80%甚至更高的准确率。训练完成后,我们可以绘制学习曲线,直观地观察训练过程。

# 绘制训练历史
def plot_training_history(history):
    acc = history.history['accuracy']
    val_acc = history.history['val_accuracy']
    loss = history.history['loss']
    val_loss = history.history['val_loss']
    epochs_range = range(len(acc))

    plt.figure(figsize=(12, 4))
    plt.subplot(1, 2, 1)
    plt.plot(epochs_range, acc, label='训练准确率')
    plt.plot(epochs_range, val_acc, label='验证准确率')
    plt.legend(loc='lower right')
    plt.title('训练和验证准确率')

    plt.subplot(1, 2, 2)
    plt.plot(epochs_range, loss, label='训练损失')
    plt.plot(epochs_range, val_loss, label='验证损失')
    plt.legend(loc='upper right')
    plt.title('训练和验证损失')
    plt.show()

plot_training_history(history)

如果验证准确率已经达到你的预期,并且没有明显的过拟合(即验证损失没有持续上升),那么模型可能已经足够好了。但如果性能还有提升空间,或者你希望模型更好地适应你数据集的独特分布,可以进入下一阶段:微调。

5. 第二阶段:微调预训练模型

微调是指解冻基础模型的一部分层,让它们也参与训练,与新分类头一起进行端到端的优化。

5.1 解冻部分层并进行微调

一个常见的策略是解冻基础模型靠顶部的若干层,因为这些层学习的是更具体、更高级的特征,而底层的通用特征(如边缘)则可以保持冻结。

# 解冻基础模型顶部的部分层
base_model.trainable = True

# 查看基础模型总层数
print(f"基础模型总层数: {len(base_model.layers)}")

# 选择解冻顶部多少层?一个经验法则是解冻最后1/4到1/3的层
fine_tune_at = int(len(base_model.layers) * 0.75) # 解冻最后25%的层
print(f"将从第 {fine_tune_at} 层开始解冻。")

# 冻结fine_tune_at之前的所有层
for layer in base_model.layers[:fine_tune_at]:
    layer.trainable = False

# 重新编译模型
# 微调时需要更小的学习率,通常比第一阶段小10倍
model.compile(
    optimizer=optimizers.Adam(learning_rate=initial_learning_rate / 10),
    loss='categorical_crossentropy',
    metrics=['accuracy']
)

# 再次查看模型摘要,确认可训练参数数量
model.summary()

在模型摘要中,你会看到“Trainable params”的数量大幅增加,这正是因为我们解冻了基础模型的部分层。

5.2 继续训练(微调)

现在,我们以更小的学习率继续训练模型。

# 设置微调轮次
fine_tune_epochs = 10
total_epochs = initial_epochs + fine_tune_epochs

# 继续训练
print("开始第二阶段训练(微调)...")
history_fine = model.fit(
    train_ds,
    initial_epoch=history.epoch[-1] + 1,
    epochs=total_epochs,
    validation_data=val_ds,
    callbacks=[checkpoint_cb, early_stopping_cb, reduce_lr_cb]
)

将两个阶段的训练历史合并,可以更完整地观察整个过程。

# 合并历史记录
acc = history.history['accuracy'] + history_fine.history['accuracy']
val_acc = history.history['val_accuracy'] + history_fine.history['val_accuracy']
loss = history.history['loss'] + history_fine.history['loss']
val_loss = history.history['val_loss'] + history_fine.history['val_loss']

epochs_range = range(total_epochs)

plt.figure(figsize=(12, 4))
plt.subplot(1, 2, 1)
plt.plot(epochs_range, acc, label='训练准确率')
plt.plot(epochs_range, val_acc, label='验证准确率')
# 标记微调开始点
plt.axvline(x=initial_epochs-1, color='gray', linestyle='--')
plt.text(initial_epochs-1, max(val_acc), ' 微调开始', color='gray')
plt.legend(loc='lower right')
plt.title('完整训练和验证准确率')

plt.subplot(1, 2, 2)
plt.plot(epochs_range, loss, label='训练损失')
plt.plot(epochs_range, val_loss, label='验证损失')
plt.axvline(x=initial_epochs-1, color='gray', linestyle='--')
plt.legend(loc='upper right')
plt.title('完整训练和验证损失')
plt.show()

理想情况下,你会看到在微调开始后,验证准确率有一个明显的、稳定的提升,而验证损失则继续下降或保持平稳。

6. 模型评估与预测

训练完成后,让我们在测试集(或保留的验证集)上最终评估模型性能,并看看它如何对单张图片进行预测。

6.1 加载最佳模型并评估

# 加载第一阶段保存的最佳模型(或根据你的需求加载最终模型)
best_model = keras.models.load_model("best_model_phase1.keras")

# 假设我们有一个未参与训练/验证的测试集
test_dir = data_dir / 'test'
test_ds = preprocessing.image_dataset_from_directory(
    test_dir,
    image_size=(IMG_HEIGHT, IMG_WIDTH),
    batch_size=BATCH_SIZE,
    label_mode='categorical',
    shuffle=False # 测试集无需打乱
)
test_ds = prepare_dataset(test_ds) # 应用相同的预处理

# 评估模型
print("在测试集上评估模型...")
test_loss, test_accuracy = best_model.evaluate(test_ds)
print(f"测试损失: {test_loss:.4f}")
print(f"测试准确率: {test_accuracy:.4f}")

6.2 单张图片预测与可视化

让我们写一个实用的函数,可以输入任意图片路径,让模型进行预测并可视化结果。

def predict_and_plot_image(model, img_path, class_names):
    """
    加载单张图片,进行预测,并显示图片与预测结果。
    """
    img = keras.utils.load_img(img_path, target_size=(IMG_HEIGHT, IMG_WIDTH))
    img_array = keras.utils.img_to_array(img)
    img_array = tf.expand_dims(img_array, 0) # 创建批次维度 (1, H, W, C)
    img_array = img_array / 255.0 # 归一化,与训练时一致

    predictions = model.predict(img_array, verbose=0)
    score = tf.nn.softmax(predictions[0]) # 获取概率分布

    plt.figure(figsize=(6, 6))
    plt.imshow(img)
    plt.axis('off')
    
    predicted_class = class_names[np.argmax(score)]
    confidence = 100 * np.max(score)
    
    plt.title(f"预测: {predicted_class}\n置信度: {confidence:.2f}%")
    plt.show()
    
    # 打印所有类别的概率
    print("各类别概率:")
    for i in range(len(class_names)):
        print(f"  {class_names[i]}: {100*score[i]:.2f}%")

# 使用示例
# 请将 `example_image.jpg` 替换为你的测试图片路径
# predict_and_plot_image(best_model, "./example_image.jpg", class_names)

6.3 生成混淆矩阵(进阶)

对于分类任务,混淆矩阵能更细致地揭示模型在哪些类别上容易混淆。

from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay
import itertools

# 获取测试集所有真实标签和预测标签
y_true = []
y_pred = []

for images, labels in test_ds:
    # labels是one-hot编码,需要转换为类别索引
    true_indices = tf.argmax(labels, axis=1).numpy()
    y_true.extend(true_indices)
    
    preds = model.predict(images, verbose=0)
    pred_indices = tf.argmax(preds, axis=1).numpy()
    y_pred.extend(pred_indices)

# 计算混淆矩阵
cm = confusion_matrix(y_true, y_pred)

# 绘制混淆矩阵
disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=class_names)
fig, ax = plt.subplots(figsize=(8, 8))
disp.plot(ax=ax, cmap=plt.cm.Blues, values_format='d')
plt.xticks(rotation=45)
plt.title("混淆矩阵")
plt.show()

通过混淆矩阵,你可以清晰地看到模型是否对某个特定类别识别率低,或者经常将A类误判为B类,这为进一步的数据收集或模型调整提供了明确方向。

7. 实战技巧与避坑指南

在项目实践中,以下几个点常常决定了迁移学习的最终成败:

1. 预训练模型的选择 并非越大的模型越好。你需要权衡精度、速度和模型大小。下表对比了几种常见模型的特点:

模型参数量(约)ImageNet Top-1 精度特点适用场景
MobileNetV23.4M71.3%极轻量,速度快移动端、嵌入式设备
EfficientNetB05.3M77.1%精度/效率平衡极佳通用服务器/云端部署
ResNet5025.6M74.9%经典、稳定、社区支持好研究、对稳定性要求高的生产环境
EfficientNetB419M82.9%精度高,但计算量大对精度要求极高的任务,算力充足

2. 学习率策略

  • 分阶段训练:如本文所示,先冻结训练顶层(较大学习率,如1e-3),再微调(较小学习率,如1e-4, 1e-5)。
  • 学习率预热:在训练开始时,从一个很小的值线性增加到初始学习率,有助于稳定训练初期。
  • 余弦退火:使用CosineDecay调度器,让学习率像余弦曲线一样平滑下降,通常比阶梯下降效果更好。

3. 过拟合的应对 小数据集是过拟合的温床。除了数据增强,还可以:

  • 加大Dropout比率:在分类头中尝试0.5甚至更高的Dropout。
  • 权重衰减:在优化器中加入L2正则化(tf.keras.optimizers.Adam(weight_decay=1e-4))。
  • 标签平滑:一种正则化技术,让模型对预测不那么“自信”,可以提高泛化能力。

4. 当效果不佳时 如果模型性能始终上不去,可以检查:

  • 数据质量:标注是否准确?类别是否平衡?是否有脏数据?
  • 领域差异:你的图片(如卫星图、素描)与ImageNet自然图像差异是否过大?考虑使用在更接近领域(如卫星图像数据集)上预训练的模型。
  • 微调程度:尝试解冻更多层,或者从更早的阶段就开始微调(即减少冻结层数)。

我在一个花卉分类的实际项目中,最初使用ResNet50,验证准确率卡在85%。后来切换到EfficientNetB0,并仔细调整了数据增强参数(特别是针对花卉的旋转和色彩抖动),最终在相同数据上将准确率提升到了93%。这个经历告诉我,模型选择后的“精调”过程,尤其是数据层面的处理,其重要性不亚于模型架构本身。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐