TensorFlow中如何持续获取数据集批次?有无内置函数?
从数据集持续获取批次数据的方案
首先肯定有内置工具!TensorFlow里的tf.data.Dataset就是专门用来处理这种批量数据迭代的官方方案,比自己手写get_batch要简洁、高效得多,还能自动处理数据打乱、多轮重复、多线程加载这些细节问题。
一、使用TensorFlow内置的tf.data.Dataset(最简方式)
用tf.data.Dataset的话,你可以直接把训练数据转换成数据集对象,然后链式调用几个方法就能搞定批次迭代,完全不用自己维护索引或打乱逻辑。
替换你原来的代码,示例如下(适配TensorFlow 2.x eager模式,这也是现在的主流用法):
import tensorflow as tf # 假设x_train、y_train是你的numpy格式训练数据 dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train)) # 打乱数据(buffer_size设为数据集总长度,确保充分打乱)→ 划分批次 → 支持多轮重复训练 dataset = dataset.shuffle(buffer_size=x_train.shape[0]).batch(batch_size).repeat() # 直接用for循环遍历数据集,每次拿到的就是不同的批次 for x_batch, y_batch in dataset: # TF2.x eager模式下直接调用训练步骤即可,不用session train_step(x_batch, y_batch) # 如果是兼容旧版session的写法,就转成numpy数组喂进去 # sess.run(train_step, feed_dict={x: x_batch.numpy(), y: y_batch.numpy()})
如果是TensorFlow 1.x的静态图模式,需要初始化迭代器:
# TF1.x 示例 iterator = dataset.make_initializable_iterator() x_batch, y_batch = iterator.get_next() with tf.Session() as sess: sess.run(iterator.initializer) for i in range(num_trains): try: x_b, y_b = sess.run([x_batch, y_batch]) sess.run(train_step, feed_dict={x: x_b, y: y_b}) except tf.errors.OutOfRangeError: # 一轮数据取完后,重新初始化迭代器开始下一轮 sess.run(iterator.initializer) x_b, y_b = sess.run([x_batch, y_batch]) sess.run(train_step, feed_dict={x: x_b, y: y_b})
这个方案的优势很明显:自动处理乱序、批次划分,还能轻松扩展预处理、多线程加载等功能,完全不会出现重复取同一批次的问题。
二、手动实现get_batch函数(如果不想用tf.data)
如果你一定要自己写get_batch,核心问题是要维护一个全局的索引指针,并且每轮训练结束后重新打乱数据,这样才能保证每次拿到的批次不重复。
示例代码如下:
import numpy as np # 全局变量,跟踪当前取到数据的哪个位置 current_index = 0 # 先打乱一次数据的索引,保证第一轮训练是乱序的 shuffled_indices = np.random.permutation(len(x_train)) def get_batch(x_train, y_train, batch_size): global current_index, shuffled_indices # 如果当前索引加批次大小超过数据集长度,说明一轮结束,重新打乱并重置索引 if current_index + batch_size > len(x_train): shuffled_indices = np.random.permutation(len(x_train)) current_index = 0 # 获取当前批次的索引 batch_indices = shuffled_indices[current_index:current_index+batch_size] current_index += batch_size return x_train[batch_indices], y_train[batch_indices] # 你的原循环可以直接使用这个函数 for i in range(num_trains): x_batch, y_batch = get_batch(x_train, y_train, batch_size) sess.run(train_step, feed_dict={x:x_batch, y:y_batch})
要是不想用全局变量,也可以把索引和打乱逻辑封装成类,代码更优雅:
class BatchGenerator: def __init__(self, x, y, batch_size): self.x = x self.y = y self.batch_size = batch_size self.current_index = 0 self._shuffle_indices() def _shuffle_indices(self): self.shuffled_indices = np.random.permutation(len(self.x)) def get_batch(self): if self.current_index + self.batch_size > len(self.x): self._shuffle_indices() self.current_index = 0 batch_indices = self.shuffled_indices[self.current_index:self.current_index+self.batch_size] self.current_index += self.batch_size return self.x[batch_indices], self.y[batch_indices] # 使用方式 generator = BatchGenerator(x_train, y_train, batch_size) for i in range(num_trains): x_batch, y_batch = generator.get_batch() sess.run(train_step, feed_dict={x:x_batch, y:y_batch})
总结一下,优先用TensorFlow内置的tf.data.Dataset,这是官方推荐的标准方案,省心又高效;如果一定要手动实现,核心就是维护好索引指针,每轮结束重新打乱数据。
内容的提问来源于stack exchange,提问作者Huzo
相关产品推荐
相关产品推荐

