You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在Google Colaboratory用TensorFlow VGG16构建三类分类Keras CNN及数据输入方法

问题解答

一、参考代码中train_data/valid_data的疑问

这段代码能运行的核心前提是**train_data和valid_data在这段代码执行前已经被定义完成**,参考代码只是省略了数据加载的环节。在Keras中,model.fit()要求输入的数据集必须是张量、Numpy数组,或是tf.data.Dataset、ImageDataGenerator生成的迭代器——这些数据对象需要你提前通过数据加载代码创建,比如从本地/Google Drive读取图像、预处理后封装成对应格式。

二、三类分类VGG16模型实现(含数据输入方案)

以下是在Google Colab上搭建三类分类VGG16模型的完整流程,覆盖数据输入、模型构建、训练全环节:

1. 数据准备(两种常用方案)

方案1:用ImageDataGenerator加载文件夹结构数据集

如果你的数据集按如下结构存放(可上传至Colab或挂载Google Drive):

dataset/
├── train/
│   ├── class_a/
│   ├── class_b/
│   └── class_c/
└── valid/
    ├── class_a/
    ├── class_b/
    └── class_c/

可通过ImageDataGenerator自动生成训练/验证数据:

from tensorflow.keras.preprocessing.image import ImageDataGenerator
from tensorflow.keras.applications.vgg16 import preprocess_input

# 训练集数据增强+VGG16专属预处理
train_datagen = ImageDataGenerator(
    preprocessing_function=preprocess_input,
    rotation_range=20,
    width_shift_range=0.2,
    height_shift_range=0.2,
    horizontal_flip=True
)

# 验证集仅做预处理
valid_datagen = ImageDataGenerator(preprocessing_function=preprocess_input)

# 加载训练数据
train_data = train_datagen.flow_from_directory(
    'dataset/train',
    target_size=(224, 224),  # 匹配VGG16输入尺寸
    batch_size=32,
    class_mode='categorical'  # 三类分类用独热编码标签
)

# 加载验证数据
valid_data = valid_datagen.flow_from_directory(
    'dataset/valid',
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical'
)

方案2:用tf.data.Dataset加载(自定义性更强)

如果需要自定义数据读取逻辑,比如从文件列表加载:

import tensorflow as tf
from tensorflow.keras.applications.vgg16 import preprocess_input

# 定义图像加载与预处理函数
def load_and_preprocess(file_path):
    # 从路径提取标签(需根据你的数据集路径规则调整)
    label = tf.strings.split(file_path, '/')[-2]
    label_map = {'class_a':0, 'class_b':1, 'class_c':2}
    label = tf.cast(tf.equal(label, list(label_map.keys())), tf.float32)
    
    # 加载并预处理图像
    img = tf.io.read_file(file_path)
    img = tf.image.decode_jpeg(img, channels=3)
    img = tf.image.resize(img, (224, 224))
    img = preprocess_input(img)
    return img, label

# 创建训练/验证数据集
train_files = tf.data.Dataset.list_files('dataset/train/*/*.jpg')
train_data = train_files.map(load_and_preprocess).batch(32).shuffle(1000)

valid_files = tf.data.Dataset.list_files('dataset/valid/*/*.jpg')
valid_data = valid_files.map(load_and_preprocess).batch(32)

2. 构建三类分类VGG16模型

调整原参考代码的输出层与损失函数,适配三类分类需求:

from tensorflow.keras.applications.vgg16 import VGG16
from tensorflow.keras import layers, models

# 加载预训练VGG16,移除顶层全连接层
base_model = VGG16(weights="imagenet", include_top=False, input_shape=(224,224,3))
base_model.trainable = False  # 先冻结预训练层

# 添加自定义分类头
model = models.Sequential([
    base_model,
    layers.Flatten(),
    layers.Dense(50, activation='relu'),
    layers.Dense(20, activation='relu'),
    layers.Dense(3, activation='softmax')  # 三类分类用softmax输出概率分布
])

# 编译模型:对应独热标签用CategoricalCrossentropy
model.compile(
    optimizer='adam',
    loss=tf.keras.losses.CategoricalCrossentropy(),
    metrics=['accuracy']
)

如果你的标签是整数形式(如0、1、2),可将class_mode设为'sparse',损失函数改用SparseCategoricalCrossentropy(),输出层仍保留softmax。

3. 训练模型

提前定义回调函数后即可启动训练:

# 定义回调示例:早停+最优模型保存
es = tf.keras.callbacks.EarlyStopping(patience=5, restore_best_weights=True)
model_cp = tf.keras.callbacks.ModelCheckpoint('best_model.h5', save_best_only=True)

# 启动训练
history = model.fit(
    train_data,
    validation_data=valid_data,
    epochs=30,
    verbose=1,
    callbacks=[es, model_cp]
)

内容的提问来源于stack exchange,提问作者user22247469

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 11:35:05