如何拆分Python generator的类别标签实现PlantVillage数据集多任务分类
解决方案
核心问题说明
- 内存崩溃的根本原因是你尝试将整个数据集一次性加载到内存中,PlantVillage数据集数万张255*255的彩色图片总大小远超普通运行内存,完全没有必要全量加载。
- 原有逻辑还存在隐藏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
相关产品推荐
相关产品推荐

