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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 03:12:48