2.2 从内存到生成器:三种上料方式


文档摘要

2.2 从内存到生成器:三种上料方式 本节摘要:构造 Dataset 有三条主路:fromtensorslices 直接共享内存数组、fromtensors 整块打包、fromgenerator 逐项产出。本节讲清三者的内存模型与适用边界——数组共享视图零拷贝、生成器能流式喂无限数据但状态要收好——并用加州房价与合成的无限流各演示一遍。上料方式选错,轻则内存翻倍,重则一张批次的形状都出不来。 本节能力清单 阅读完本节,你应当能够: 说清 fromtensorslices 与 fromtensors 的切片语义差异,现场预测输出形状; 编写带 outputsignature 的生成器管道,实现无限流训练数据; 根据数据规模与产生方式正确选择三种上料方式;

2.2 从内存到生成器:三种上料方式

本节摘要:构造 Dataset 有三条主路:from_tensor_slices 直接共享内存数组、from_tensors 整块打包、from_generator 逐项产出。本节讲清三者的内存模型与适用边界——数组共享视图零拷贝、生成器能流式喂无限数据但状态要收好——并用加州房价与合成的无限流各演示一遍。上料方式选错,轻则内存翻倍,重则一张批次的形状都出不来。

本节能力清单

阅读完本节,你应当能够:

  1. 说清 from_tensor_slices 与 from_tensors 的切片语义差异,现场预测输出形状;
  2. 编写带 output_signature 的生成器管道,实现无限流训练数据;
  3. 根据数据规模与产生方式正确选择三种上料方式;
  4. 排查生成器管道"形状未知、类型不匹配"两类报错。

三个入口,三种内存观

from_tensor_slices 接收数组或数组元组,沿第 0 维切成一个个元素:传入形状 (20640, 8) 的特征矩阵,得到 20640 个形状 (8,) 的样本。它不复制数据——底层数组被共享为视图,内存占用近似为零。from_tensors 则把整个输入当成一个元素打包:传进去什么,数据集里就只有一个"整体张量"。from_generator 接一个 Python 生成器函数,每次迭代调用它产出下一批元素,数据可以根本不存在内存里——按需现做。

import tensorflow as tf import numpy as np a = tf.constant([[1, 2], [3, 4], [5, 6]]) # (3, 2) slices = tf.data.Dataset.from_tensor_slices(a) print(list(slices.as_numpy_iterator())) # 输出:[array([1, 2]), array([3, 4]), array([5, 6])] # 沿轴 0 切成 3 个样本,每个 shape (2,) tensors = tf.data.Dataset.from_tensors(a) print(list(tensors.as_numpy_iterator())) # 输出:[array([[1, 2], # [3, 4], # [5, 6]])] # 整块作为一个元素,数据集长度为 1 print(tf.data.Dataset.cardinality(slices).numpy(), tf.data.Dataset.cardinality(tensors).numpy()) # 输出:3 1 —— 卡片计数直接暴露差异
方式 内存模型 典型场景
from_tensor_slices 共享视图,零拷贝 数组已装进内存的中小数据
from_tensors 整块入列 把单个大张量当整体喂(少见)
from_generator 按需现做,可无限 流式合成、动态生成、超大集代理

切片语义的进阶:结构化样本

from_tensor_slices 对字典、元组同样沿第 0 轴切,这对"特征与标签同行"的数据非常顺手:

features = {"面积": [120.0, 85.0, 64.0], "楼层": [3.0, 7.0, 1.0]} labels = [430.0, 300.0, 210.0] ds = tf.data.Dataset.from_tensor_slices((features, labels)) for f, y in ds.take(2): print(f, y) # 输出: # {'面积': <tf.Tensor: shape=(), dtype=float32, numpy=120.0>, # '楼层': <tf.Tensor: shape=(), dtype=float32, numpy=3.0>} 430.0 # {'面积': ..., '楼层': ...} 300.0 # 每个元素是 (特征字典, 标量标签) 对——直接对接 Keras 的字典输入模型

注意切片语义的边界:切的是第 0 轴。时间序列数据想按窗口切,直接 slices 会把时间轴切碎,得先在 NumPy 侧造好窗口再上管道——用电负荷预测的窗口构造在第 4 章 RNN 一节展开。

生成器上料:无限流与签名声明

from_generator 的核心麻烦是类型安全:生成器产出的东西,框架在跟踪时必须知道规格,所以要配 output_signature 声明。写一个合成的无限回归数据流:

def synthetic_stream(): """无限流生成器:每次 yield 一个样本对""" i = 0 while True: x = tf.random.normal([8]) # 真实关系造进标签里:y 约等于前两维的线性组合加噪声 y = tf.reduce_sum(x[:2] * 2.0) + tf.random.normal([], stddev=0.1) yield x, y i += 1 sig = (tf.TensorSpec(shape=(8,), dtype=tf.float32), tf.TensorSpec(shape=(), dtype=tf.float32)) ds = tf.data.Dataset.from_generator(synthetic_stream, output_signature=sig) ds = ds.batch(32).prefetch(tf.data.AUTOTUNE) for x_batch, y_batch in ds.take(2): print(x_batch.shape, float(tf.reduce_mean(y_batch))) # 输出: # (32, 8) -0.04(随机数,数值不定) # (32, 8) 0.21 # take 才停——生成器本身无限,管道按需拉动

生成器是 Python 函数,框架在后台线程里拉动它,与计算重叠。三条纪律:生成器必须可重入(重新迭代等于重新调用函数,所以别把状态存在函数外的可变对象里);yield 的规格必须与签名严格一致,差一个维度就抛类型错误;无限流必须有人喊停——take、steps_per_epoch 或 fit 的 epochs,三者之一。

# 纪律一的验证:重新迭代,生成器自动重开 ds2 = tf.data.Dataset.from_generator(synthetic_stream, output_signature=sig) n = 0 for _ in ds2.batch(16).take(1): n += 1 for _ in ds2.batch(16).take(1): n += 1 print("iterated twice, batches:", n) # 输出:iterated twice, batches: 2 # 第二次 take 正常工作,因为函数调用被重新发起

选择决策与排错

选型按数据来源倒推:已经是有规模上限的数组,from_tensor_slices,零成本;数据来自惰性合成或外部按需产生的逻辑,from_generator;想把"一个整体"作为单元素(比如把整段序列作为一个样本再切窗口),from_tensors。内存装不下的真实大数据别硬选生成器模拟——那是 2.3 节 TFRecord 的领地,文件格式能换来自由的随机访问与并行读取。

生成器管道两类高频报错:其一,"Dataset.from_generator 的 output_signature 与产出不符",报错信息会给出两边规格,逐维对照即可,多数是漏了 batch 前后的形状差异;其二,"生成器返回了 numpy 数组而非张量"在部分版本是允许的,但规格声明必须按最终转成的 dtype 写,建议生成器内部就 tf.convert_to_tensor 一次,把类型钉死在源头。

⚠️ 常见坑:给 from_tensor_slices 传长度不一致的特征与标签(20640 条特征、20639 条标签),报错出现在构造时就该庆幸——它至少会在管道声明阶段暴露。若长度错位发生在生成器里,问题会推迟到迭代深处,更难追。

💡 关键直觉:上料方式决定的是"元素什么时候诞生"。切片法元素早已躺在内存,生成器法元素在被拉动的瞬间诞生。凡是"训练时才做才省"的事(合成、扰动、动态采样),都应该往生成器或 map 里搬。

本节要点回顾

  • 切片与整块:from_tensor_slices 沿轴 0 切元素,from_tensors 整体单元素,卡片计数立辨。
  • 零拷贝:切片法共享底层数组视图,内存友好。
  • 结构化切片:字典元组同行切,直接对接字典输入模型。
  • 生成器三纪律:可重入、签名一致、有人喊停。
  • 大数据不硬扛:内存装不下走文件格式,别用生成器模拟文件。

下节进入文件读取,TFRecord 是大规模数据的正装。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U