How to save a tensorflow dataset to multiple shards without using enumerate

Viewed 33

I have a tensorflow dataset with some elements in it, and I want to save it with tf.data.Dataset.save such that each element gets its own shard. Thus if the dataset contains 2,000 elements, it would be saved to 2,000 shards.

The documentation here specifies how to create 1 shard only, but not how to make a shard for each element.

Below, I am able to do it with enumerate, but is there another way to do it without also saving the index from enumerate?

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)
0 Answers
Related