如何在tf.data.Dataset.from_generator中保留字典键以处理非矩形形状数据
保留字典键的tf.data.Dataset生成方案
你遇到的问题核心是:用from_generator构建数据集时,原代码输出的是元组结构,没法保留data、labels这类自定义键。其实只需要调整生成器的返回格式,把每个样本包装成字典就能解决,具体实现如下:
修改后的完整代码
import tensorflow as tf import numpy as np # 你的非矩形原始数据 samples_train = { 'data': [np.array([[1,1]]), np.array([[1,1],[2,2]]), np.array([[1,1],[2,2],[3,3]])], 'labels': [1,2,3] } # 生成带字典键的数据集 ds = tf.data.Dataset.from_generator( # 生成器返回字典,对应原始数据的键 lambda: ({'data': d, 'labels': l} for d, l in zip(samples_train['data'], samples_train['labels'])), # 输出类型也要对应成字典格式,指定每个键的数据类型 output_types={'data': tf.float32, 'labels': tf.float32} )
验证访问效果
遍历数据集时,就可以像你期望的那样通过键来访问数据了:
for batch in ds: print("data部分:") print(batch['data']) print("labels部分:") print(batch['labels']) print("---")
输出结果:
data部分: tf.Tensor([[1. 1.]], shape=(1, 2), dtype=float32) labels部分: tf.Tensor(1.0, shape=(), dtype=float32) --- data部分: tf.Tensor( [[1. 1.] [2. 2.]], shape=(2, 2), dtype=float32) labels部分: tf.Tensor(2.0, shape=(), dtype=float32) --- data部分: tf.Tensor( [[1. 1.] [2. 2.] [3. 3.]], shape=(3, 2), dtype=float32) labels部分: tf.Tensor(3.0, shape=(), dtype=float32) ---
关键细节说明
- 生成器逻辑调整:不再返回元组
(d, l),而是返回和原始数据结构一致的字典{'data': d, 'labels': l},这样每个样本都会保留你需要的键。 - 输出类型匹配:
output_types必须对应改成字典格式,明确每个键对应的数据类型,确保TensorFlow能正确解析非矩形数据。 - 扩展性:如果后续还有其他自定义键(比如
sample_id),只需要在生成器的字典里添加对应键值对,同时在output_types里补充对应类型即可。
内容的提问来源于stack exchange,提问作者Crysers
相关产品推荐
相关产品推荐

