如何将TensorFlow数据集随机划分为N个无重叠的等大小子数据集
问题分析:为啥之前的方法都失效了?
你踩的这个坑其实是TensorFlow Dataset的惰性执行机制导致的:
- 当你调用
ds.shuffle()后,这个操作并不会立刻执行,而是在每次迭代数据集的时候才会重新触发洗牌。所以你用ds.shard(N, i)的时候,每个shard在迭代时都会重新跑一遍shuffle,相当于两个shard是从两次不同的洗牌结果里取样本,自然会出现重叠,根本不是真正的划分。 - 用
take+skip的思路同理:每次take和skip都会重新触发shuffle,导致后续的skip是从新的洗牌结果里跳,完全达不到拆分的效果。
解决方案
方案1:小数据集直接转列表拆分(简单直观)
如果你的数据集不大,直接把数据转成Python列表打乱后拆分是最省心的方式:
import tensorflow as tf import math import random # 构造原始数据集 ds = tf.data.Dataset.from_tensor_slices(list(range(1, 21))) N = 2 # 转成列表并随机打乱 data_samples = list(ds.as_numpy_iterator()) random.shuffle(data_samples) # 计算每个子集的大小,拆分数据集 subset_size = math.floor(len(data_samples) / N) ds_list = [] for i in range(N): start_idx = i * subset_size # 最后一个子集要包含剩余的所有样本 end_idx = start_idx + subset_size if i != N-1 else len(data_samples) subset_data = data_samples[start_idx:end_idx] ds_list.append(tf.data.Dataset.from_tensor_slices(subset_data)) # 验证结果 for idx, sub_ds in enumerate(ds_list): sorted_samples = sorted(list(sub_ds.as_numpy_iterator())) print(f"子集{idx+1}: {sorted_samples}")
这个方法逻辑清晰、容易调试,但如果数据集太大,转成列表会占用大量内存,只适合小数据集场景。
方案2:大数据集的惰性划分(推荐)
如果你的数据集很大,没法一次性加载到内存,就用随机索引+过滤的方式,保持Dataset的惰性执行特性:
import tensorflow as tf import math # 构造原始数据集 ds = tf.data.Dataset.from_tensor_slices(list(range(1, 21))) N = 2 # 给每个样本分配一个唯一的随机索引(范围足够大保证随机性) # 先用enumerate标记原始位置(可选,避免极端情况的重复),再生成随机索引 ds = ds.enumerate().map(lambda orig_idx, x: (tf.random.uniform(shape=[], minval=0, maxval=100000, dtype=tf.int32), x)) # 根据随机索引的模N值划分数据集 ds_list = [] for i in range(N): # 筛选出随机索引模N等于当前i的样本 subset_ds = ds.filter(lambda rand_idx, x: rand_idx % N == i) # 去掉随机索引,只保留原始样本 subset_ds = subset_ds.map(lambda rand_idx, x: x) ds_list.append(subset_ds) # 验证结果 for idx, sub_ds in enumerate(ds_list): sorted_samples = sorted(list(sub_ds.as_numpy_iterator())) print(f"子集{idx+1}: {sorted_samples}")
这个方法的核心是:每个样本的随机索引只会生成一次,后续的过滤操作都是基于同一个随机索引集合,所以不会出现样本重叠,同时保持了Dataset的惰性,完美适配大数据集。
方案3:用官方的split API(TensorFlow 2.10+)
如果你的TensorFlow版本在2.10及以上,可以直接用官方提供的tf.data.experimental.split,一步搞定:
import tensorflow as tf # 构造原始数据集 ds = tf.data.Dataset.from_tensor_slices(list(range(1, 21))) N = 2 # 先洗牌,再拆分 ds = ds.shuffle(buffer_size=20) ds_list = tf.data.experimental.split(ds, num_split=N) # 验证结果 for idx, sub_ds in enumerate(ds_list): sorted_samples = sorted(list(sub_ds.as_numpy_iterator())) print(f"子集{idx+1}: {sorted_samples}")
这个API内部已经处理了惰性执行的问题,会自动将洗牌后的数据集划分为N个无重叠的子集,代码最简洁,优先推荐使用。
内容的提问来源于stack exchange,提问作者jackve
相关产品推荐
相关产品推荐

