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

tf.data.Dataset apply()方法无法更新数据集问题求助

解决PrefetchDataset标签二值化的问题

你的代码核心问题是用错了数据集转换方法,以及标签修改逻辑未生效,具体问题和修复方案如下:

问题分析

  1. apply方法用途错误:apply是对整个数据集对象执行全局转换(比如设置缓存、预取规则),不是用来逐个处理批次内容的。你把它当成批次迭代器遍历,完全不符合其设计逻辑。
  2. 标签修改未生效:即便遍历到了批次,你只是在函数内部修改了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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 13:30:42