《Keras 图像分类实战》

《Keras 图像分类实战》

用 Keras 处理图像分类任务时,最核心的优势是数据量不大也能训练出有用的模型:通过数据增强(data augmentation)与小网络配合,可以在很少数据上获得较好效果。本文介绍用 ImageDataGenerator 做数据增强 + CNN 的训练流程。

参考:https://blog.keras.io/building-powerful-image-classification-models-using-very-little-data.html

1 数据准备

典型目录结构(Keras flow_from_directory 要求):

data/ ├── train/ │ ├── cats/ # 猫图 │ └── dogs/ # 狗图 └── validation/ ├── cats/ └── dogs/

2 数据增强(ImageDataGenerator)

from tensorflow.keras.preprocessing.image import ImageDataGenerator # 训练集:随机旋转、平移、缩放、水平翻转,并归一化 train_datagen = ImageDataGenerator( rescale=1.0 / 255, rotation_range=40, width_shift_range=0.2, height_shift_range=0.2, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, ) # 验证集:只归一化,不做增强 validation_datagen = ImageDataGenerator(rescale=1.0 / 255) train_generator = train_datagen.flow_from_directory( "data/train", target_size=(150, 150), # 统一缩放 batch_size=32, class_mode="binary", # 二分类 ) validation_generator = validation_datagen.flow_from_directory( "data/validation", target_size=(150, 150), batch_size=32, class_mode="binary", )

3 小 CNN 模型

from tensorflow import keras from tensorflow.keras import layers model = keras.Sequential([ layers.Conv2D(32, (3, 3), activation="relu", input_shape=(150, 150, 3)), layers.MaxPooling2D(2, 2), layers.Conv2D(64, (3, 3), activation="relu"), layers.MaxPooling2D(2, 2), layers.Conv2D(128, (3, 3), activation="relu"), layers.MaxPooling2D(2, 2), layers.Conv2D(128, (3, 3), activation="relu"), layers.MaxPooling2D(2, 2), layers.Flatten(), layers.Dropout(0.5), layers.Dense(512, activation="relu"), layers.Dense(1, activation="sigmoid"), # 二分类 ]) model.compile( optimizer=keras.optimizers.Adam(learning_rate=1e-4), loss="binary_crossentropy", metrics=["accuracy"], ) model.summary()

4 训练与保存

history = model.fit( train_generator, steps_per_epoch=100, # 每 epoch 的批数 epochs=30, validation_data=validation_generator, validation_steps=50, ) model.save("cats_vs_dogs.h5")

5 数据量更少时的进阶技巧

  • 迁移学习(推荐):使用预训练网络(VGG16 / ResNet / EfficientNet)做特征提取或微调,小数据集下收敛更快、效果更好:
from tensorflow.keras.applications import VGG16 base = VGG16(weights="imagenet", include_top=False, input_shape=(150, 150, 3)) base.trainable = False # 冻结主干 model = keras.Sequential([ base, layers.Flatten(), layers.Dropout(0.5), layers.Dense(1, activation="sigmoid"), ])
  • 学习率衰减:使用 ReduceLROnPlateau 回调,在验证 loss 平台期自动降低学习率:
keras.callbacks.ReduceLROnPlateau(monitor="val_loss", factor=0.5, patience=5)

6 常见问题

  • 过拟合明显(训练 acc 高、验证 acc 低):增加数据增强强度、增大 Dropout、冻结更多层或减小网络。
  • flow_from_directory 路径报错:确认子目录名称即类别名,目录下必须有图片。
  • 显存不足:调小 target_size(如 128x128)或 batch_size。

参考文档

阅读 — · 全站 —
🎸 我的歌单 0 首