如何不使用enumerate将TensorFlow数据集保存为多个独立分片
问题描述
- 需求:调用
tf.data.Dataset.save保存数据集时,实现每个元素单独存储为一个分片,例如数据集包含2000个元素时,保存后生成2000个独立分片文件。 - 现有公开文档仅覆盖单分片、固定数量分片的保存方法,未说明单元素独立分片的实现方式。
- 当前可通过
enumerate拼接索引+自定义shard_func实现需求,但该方案会将enumerate生成的索引作为元素的一部分存入文件,需要找到不额外存储索引的实现方案。
当前可运行的带索引实现代码如下:
import numpy as np import tensorflow as tf tuple_data = np.array([3, 4]) data = tf.data.Dataset.from_tensor_slices(tuple_data) data = data.enumerate() print(list(data.as_numpy_iterator())) # [(0, 3), (1, 4)] data.save(path='~/Desktop/1', shard_func=lambda i, x: i)
实现方案
不需要修改数据集本身的元素结构,通过外部维护自增计数器的方式即可实现需求,保存的文件不会包含额外索引字段,加载后和原始数据集结构完全一致。该方案不依赖数据集本身的元素结构,无论元素是单张量、元组还是字典格式,都可以正常使用。
import os import numpy as np import tensorflow as tf # 构建原始测试数据集 tuple_data = np.array([3, 4]) data = tf.data.Dataset.from_tensor_slices(tuple_data) # 初始化分片ID计数器,类型用tf.int64匹配tf.data内部要求 shard_counter = tf.Variable(0, dtype=tf.int64, trainable=False) def shard_func(_): current_id = shard_counter.read_value() shard_counter.assign_add(1) return current_id # 注意路径中的~需要手动展开,tf.io不支持自动解析~ save_path = os.path.expanduser('~/Desktop/1') data.save(save_path, shard_func=shard_func) # 验证加载结果 loaded_ds = tf.data.Dataset.load(save_path) print(list(loaded_ds.as_numpy_iterator())) # 输出: [3, 4],无额外索引字段,和原始数据集结构完全一致
注意事项
- 保存前不要对数据集执行会触发全量遍历的操作(例如提前调用
list(data.as_numpy_iterator())),否则会导致计数器提前自增,出现分片ID错位。 - 该方案生成的分片数和数据集元素总数完全相等,符合单元素单分片的需求。
- 如果需要重复保存,每次保存前需要将计数器重置为0,避免分片ID重复。
内容的提问来源于stack exchange,提问作者Jesse Kerr
相关产品推荐
相关产品推荐

