在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
相关产品推荐
相关产品推荐

