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

如何拆分Python generator的类别标签实现PlantVillage数据集多任务分类

解决方案

核心问题说明

  1. 内存崩溃的根本原因是你尝试将整个数据集一次性加载到内存中,PlantVillage数据集数万张255*255的彩色图片总大小远超普通运行内存,完全没有必要全量加载。
  2. 原有逻辑还存在隐藏bug:image_dataset_from_directory设置label_mode='categorical'返回的标签是one-hot编码的向量,不是原始的类名字符串,后续split('___')操作根本无法执行。

最优实现方案(无需自定义生成器,基于tf.data流水线处理)

该方案利用TensorFlow自带的数据集流水线逐批次加载、处理数据,不会占用过量内存,运行效率远高于自定义生成器。

第一步:加载数据集,获取类名映射

import tensorflow as tf
BATCH_SIZE = 32
IMG_SIZE = (255, 255)
data_dir = "/content/plantvillage dataset/color"

# 加载数据集时label_mode用int,返回原始类别索引
train_dataset = tf.keras.utils.image_dataset_from_directory(
    data_dir,
    shuffle=True,
    label_mode = 'int',
    validation_split = 0.2,
    batch_size=BATCH_SIZE,
    seed = 42,
    subset = "training",
    image_size=IMG_SIZE
)
validation_dataset = tf.keras.utils.image_dataset_from_directory(
    data_dir,
    shuffle=True,
    label_mode = 'int',
    validation_split = 0.2,
    batch_size=BATCH_SIZE,
    seed = 42,
    subset = "validation",
    image_size=IMG_SIZE
)

第二步:提前构建标签映射表

提前把所有类名拆分,生成原始类别索引到物种标签、病害标签的映射,不需要每次处理样本都重复拆分字符串:

# 获取所有原始类名
class_names = train_dataset.class_names

# 生成物种、病害的编码表
specie_set = set()
disease_set = set()
for name in class_names:
    s, d = name.split('___')
    specie_set.add(s)
    disease_set.add(d)
specie_to_idx = {s:i for i,s in enumerate(specie_set)}
disease_to_idx = {d:i for i,d in enumerate(disease_set)}
NUM_SPECIES = len(specie_set)
NUM_DISEASES = len(disease_set)

# 构建 原始类别id -> (物种id, 病害id) 的映射表
label_mapping = {}
for class_id, class_name in enumerate(class_names):
    s, d = class_name.split('___')
    label_mapping[class_id] = (specie_to_idx[s], disease_to_idx[d])

第三步:用map函数逐批次处理标签

tf.data的流水线会逐批次加载、处理数据,不会一次性把所有数据加载到内存,完美解决OOM问题:

def process_single_sample(image, label):
    # 从映射表取出两个标签
    def get_multi_labels(class_id):
        return label_mapping[int(class_id)]
    specie_label, disease_label = tf.numpy_function(
        get_multi_labels, 
        inp=[label], 
        Tout=[tf.int32, tf.int32]
    )
    # 固定标签shape,避免后续训练报错
    specie_label.set_shape(())
    disease_label.set_shape(())
    # 如果需要one-hot编码,取消下面两行注释即可
    # specie_label = tf.one_hot(specie_label, NUM_SPECIES)
    # disease_label = tf.one_hot(disease_label, NUM_DISEASES)
    return image, (specie_label, disease_label)

# 给训练、验证集应用处理逻辑,开多线程加速+预加载
train_dataset = train_dataset.map(
    process_single_sample, 
    num_parallel_calls=tf.data.AUTOTUNE
).prefetch(tf.data.AUTOTUNE)
validation_dataset = validation_dataset.map(
    process_single_sample, 
    num_parallel_calls=tf.data.AUTOTUNE
).prefetch(tf.data.AUTOTUNE)

后续训练示例

处理后的数据集可以直接传入model.fit(),只需构建对应两个输出头的多任务模型即可:

inputs = tf.keras.Input(shape=(255,255,3))
backbone = tf.keras.applications.MobileNetV2(include_top=False, input_tensor=inputs)
x = tf.keras.layers.GlobalAveragePooling2D()(backbone.output)
# 物种分类头
specie_output = tf.keras.layers.Dense(NUM_SPECIES, activation='softmax', name='specie_cls')(x)
# 病害分类头
disease_output = tf.keras.layers.Dense(NUM_DISEASES, activation='softmax', name='disease_cls')(x)
model = tf.keras.Model(inputs=inputs, outputs=[specie_output, disease_output])

# 编译模型,分别给两个头设置损失、权重
model.compile(
    optimizer='adam',
    loss={
        'specie_cls': 'sparse_categorical_crossentropy',
        'disease_cls': 'sparse_categorical_crossentropy'
    },
    metrics=['accuracy']
)
# 直接传入处理好的数据集训练
model.fit(train_dataset, validation_data=validation_dataset, epochs=10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 11:15:08