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

如何在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)
---

关键细节说明

  1. 生成器逻辑调整:不再返回元组(d, l),而是返回和原始数据结构一致的字典{'data': d, 'labels': l},这样每个样本都会保留你需要的键。
  2. 输出类型匹配:output_types必须对应改成字典格式,明确每个键对应的数据类型,确保TensorFlow能正确解析非矩形数据。
  3. 扩展性:如果后续还有其他自定义键(比如sample_id),只需要在生成器的字典里添加对应键值对,同时在output_types里补充对应类型即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 04:17:33