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

如何从TensorFlow Dataset手动截取前N个样本并保留顺序?

手动拆分MNIST训练集为训练/验证集(保留原始顺序)

你现有生成TensorFlow Dataset迭代器的代码:

@tf.function
def normalize_image(record):
  out = record.copy()
  out['image'] = tf.cast(out['image'], 'float32') / 255.
  return out


train_it = iter(tfds.builder('mnist').as_dataset(split='train').map(normalize_image).repeat().batch(256*10))

需要将MNIST的60000条训练样本手动拆分:前50000条作为训练集,剩余10000条作为验证集,且要保留样本原始顺序,同时能继续使用Dataset的map等操作。你尝试转为NumPy数组拆分后无法再用map,想过转存PNG再加载但不确定顺序。

解决方案:直接用TFDS的Split切片语法

TFDS原生支持对数据集做切片拆分,无需转NumPy,既能严格保留原始顺序,又能继续使用Dataset的所有操作。

完整代码示例:

import tensorflow as tf
import tensorflow_datasets as tfds

@tf.function
def normalize_image(record):
  out = record.copy()
  out['image'] = tf.cast(out['image'], 'float32') / 255.
  return out

# 加载前50000条作为训练集
train_ds = tfds.builder('mnist').as_dataset(split='train[:50000]')
train_ds = train_ds.map(normalize_image).repeat().batch(256*10)
train_it = iter(train_ds)

# 加载剩余10000条作为验证集
val_ds = tfds.builder('mnist').as_dataset(split='train[50000:]')
val_ds = val_ds.map(normalize_image).batch(256*10)
val_it = iter(val_ds)

说明

  • TFDS的MNIST训练集严格按照原始数据顺序存储,train[:50000]会精确取前50000条样本,train[50000:]取后10000条,完全符合你的顺序要求。
  • 这种方式不需要转换为NumPy数组,因此可以正常使用map、batch、repeat等Dataset操作,保留TensorFlow的数据流优势。
  • 无需额外转存PNG,避免不必要的磁盘IO和数据格式转换风险。

内容的提问来源于stack exchange,提问作者m0ss

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 12:00:54