如何高效为TensorFlow数据集中的每个y添加唯一ID值
你之前用全局变量的方案失效,是因为tf.data.Dataset.map默认在TensorFlow图模式下执行,Python层面的全局变量不会在图的每次迭代中被更新,所以所有样本拿到的id都是初始值。
你需要的功能可以用TensorFlow原生的tf.data.Dataset.enumerate()实现,这个是内置的并行化操作,全程在TensorFlow图中执行,完全不会有手动遍历的性能问题,你之前觉得枚举耗时极长,大概率是用了Python侧循环遍历数据集的错误实现方式。
实现代码
方案1:id随迭代顺序生成(shuffle后id重新分配)
如果不需要id和固定样本绑定,只需要每次迭代时每个样本拿到唯一id,直接在shuffle之后调用enumerate即可:
(ds_train_original, ds_test_original), ds_info = tfds.load( "mnist", split=["train", "test"], shuffle_files=True, as_supervised=True, with_info=True, ) batch_size = 2048 def normalize_img(image, label): """Normalizes images: `uint8` -> `float32`.""" return tf.cast(image, tf.float32) / 255.0, label # 归一化、缓存、shuffle逻辑和你原有逻辑一致 ds_train_original = ds_train_original.map(normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_train = ds_train_original.cache() ds_train = ds_train.shuffle(ds_info.splits["train"].num_examples) # 生成唯一自增id,调整为你需要的(x, y)结构 ds_train = ds_train.enumerate().map( lambda idx, (img, _): (img, (img, idx)), num_parallel_calls=tf.data.AUTOTUNE ) # 后续正常batch、prefetch即可 ds_train = ds_train.batch(batch_size).prefetch(tf.data.AUTOTUNE)
方案2:id和样本绑定(shuffle后id不变)
如果需要同一个图片每次拿到的id都是固定值,不随shuffle顺序变化,调整enumerate的调用顺序,在shuffle之前完成id和样本的绑定即可:
ds_train_original = ds_train_original.map(normalize_img, num_parallel_calls=tf.data.AUTOTUNE) # 先给每个样本绑定固定id ds_train = ds_train_original.enumerate().map( lambda idx, (img, _): (img, idx), num_parallel_calls=tf.data.AUTOTUNE ) ds_train = ds_train.cache() # 再执行shuffle,id已经和样本绑定不会变化 ds_train = ds_train.shuffle(ds_info.splits["train"].num_examples) # 调整为目标结构 ds_train = ds_train.map( lambda img, idx: (img, (img, idx)), num_parallel_calls=tf.data.AUTOTUNE )
两种方案都是纯TensorFlow原生操作,没有Python侧的遍历逻辑,性能和原生数据集加载几乎没有差异。
内容的提问来源于stack exchange,提问作者Malcolm
相关产品推荐
相关产品推荐

