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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 03:57:26