tf.data.Dataset apply()方法无法更新数据集问题求助
解决PrefetchDataset标签二值化的问题
你的代码核心问题是用错了数据集转换方法,以及标签修改逻辑未生效,具体问题和修复方案如下:
问题分析
apply方法用途错误:apply是对整个数据集对象执行全局转换(比如设置缓存、预取规则),不是用来逐个处理批次内容的。你把它当成批次迭代器遍历,完全不符合其设计逻辑。- 标签修改未生效:即便遍历到了批次,你只是在函数内部修改了
labels变量,但没有返回修改后的(图像, 新标签)对,原数据集的批次根本没被更新。
正确实现方案
用map方法替代apply,map会自动遍历数据集的每个(images, labels)批次,执行转换逻辑并返回新的批次:
batch_size = 32 img_height = 250 img_width = 250 train_ds = image_dataset_from_directory( data_dir, validation_split=0.2, color_mode="rgb", subset="training", seed=69, crop_to_aspect_ratio=False, image_size=(img_height, img_width), batch_size=batch_size) class_names = train_ds.class_names # ['Painting', 'Photo', 'Schematics', 'Sketch', 'Text'] # 通过类名动态获取索引,避免硬编码 photo_class_idx = class_names.index('Photo') def convert_to_binary_label(images, labels): # 使用TensorFlow原生操作,比Python循环效率更高,且兼容图模式 binary_labels = tf.cast(tf.equal(labels, photo_class_idx), tf.int32) return images, binary_labels # 用map方法执行标签转换 new_train_ds = train_ds.map(convert_to_binary_label)
验证转换效果
可以取一个批次查看标签是否正确转换:
for images, labels in new_train_ds.take(1): print(labels.numpy()) # 输出只会包含0和1
内容的提问来源于stack exchange,提问作者Eccsx
相关产品推荐
相关产品推荐

