如何将Kaggle自定义数据集适配Keras Vision Transformer示例代码
将自定义图像数据集适配到Keras Vision Transformer图像分类示例
以下是将ASL字母数据集适配到Keras官方Vision Transformer(ViT)示例的具体步骤:
1. 数据集预处理
ASL字母数据集采用按类别分文件夹的结构(每个文件夹对应一个手势/字母,文件夹名为类别标签)。先完成以下准备:
- 下载并解压数据集到本地目录
- 若数据集未划分训练/验证集,可通过代码在加载时自动拆分(无需手动移动文件)
2. 替换数据加载逻辑
原示例使用tf.keras.datasets.cifar100加载内置数据集,我们改用Keras的image_dataset_from_directory加载自定义文件夹结构的数据集:
import tensorflow as tf # 定义图像参数(ASL图像默认尺寸为200x200,需替换原示例的32x32) IMAGE_SIZE = (200, 200) BATCH_SIZE = 32 # 加载训练集与验证集(自动按20%比例拆分验证集) train_ds = tf.keras.utils.image_dataset_from_directory( "本地数据集路径/asl_alphabet_train", validation_split=0.2, subset="training", seed=123, image_size=IMAGE_SIZE, batch_size=BATCH_SIZE, ) val_ds = tf.keras.utils.image_dataset_from_directory( "本地数据集路径/asl_alphabet_train", validation_split=0.2, subset="validation", seed=123, image_size=IMAGE_SIZE, batch_size=BATCH_SIZE, ) # 图像归一化到[0,1]范围 normalization_layer = tf.keras.layers.Rescaling(1./255) train_ds = train_ds.map(lambda x, y: (normalization_layer(x), y)) val_ds = val_ds.map(lambda x, y: (normalization_layer(x), y)) # 缓存+预取提升训练性能 AUTOTUNE = tf.data.AUTOTUNE train_ds = train_ds.cache().prefetch(buffer_size=AUTOTUNE) val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)
3. 调整ViT核心参数
原示例针对32x32的CIFAR100图像设计,需根据ASL图像尺寸和类别数修改关键参数:
# 调整patch尺寸:200x200图像拆分为8x8个25x25的patch patch_size = (25, 25) num_patches = (IMAGE_SIZE[0] // patch_size[0]) ** 2 # 模型容量参数可按需微调,以下为适配ASL的参考值 projection_dim = 64 num_heads = 4 transformer_units = [projection_dim * 2, projection_dim] transformer_layers = 8 mlp_head_units = [2048, 1024] # 替换类别数:ASL数据集共29个类别(26字母+3特殊手势) num_classes = 29
4. 可选添加数据增强
针对手势图像的特点,可添加数据增强提升模型泛化能力:
data_augmentation = tf.keras.Sequential( [ tf.keras.layers.RandomFlip("horizontal"), tf.keras.layers.RandomRotation(0.1), tf.keras.layers.RandomZoom(0.1), ] ) # 将增强逻辑应用到训练集 train_ds = train_ds.map(lambda x, y: (data_augmentation(x, training=True), y))
5. 训练与评估
保留原示例的优化器、学习率调度器和训练逻辑,仅需将训练轮次epochs调整为适配ASL数据集的数值(建议20-50轮):
# 保留原示例的模型构建、优化器、学习率调度器代码 # ... # 启动训练 history = model.fit( train_ds, validation_data=val_ds, epochs=30 # 根据训练效果调整 ) # 模型评估 model.evaluate(val_ds)
内容的提问来源于stack exchange,提问作者jason hartanto
相关产品推荐
相关产品推荐

