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

如何基于含不等长列表的字典创建TF Dataset并解决转换报错

解决方案

报错本质是tf.data.Dataset.from_tensor_slices默认会将输入序列转换为规则张量,而示例中a、b两个字段包含长度不一致的数组,无法构成规则张量。直接对整个字典调用tf.ragged.constant无法生效,是因为该方法会将整个字典作为单序列处理,而非按字典内的字段分别处理,以下是两种可行实现方案:

方案1:单独转换可变长字段为RaggedTensor

这是最轻量化的方案,无需改动原有数据逻辑,仅对字典内的可变长字段单独做不规则张量转换即可:

import tensorflow as tf
import numpy as np

t_dic = {
    "uuid": np.array(["abc", "def", "ghi", "pqr"]),
    # 可变长字段单独转为不规则张量
    "a": tf.ragged.constant([np.array([1, 2, 3]), 
                             np.array([6, 2, 3]), 
                             np.array([6, 8, 1]), 
                             np.array([6, 2, 3, 10])]),
    "b": tf.ragged.constant([np.array(["a", "f", "f"]), 
                             np.array(["aa", "ff", "fs"]), 
                             np.array(["aa", "ff", "fs"]), 
                             np.array(["aa", "ff", "fs", "ss"])])
}

x = tf.data.Dataset.from_tensor_slices(t_dic)

可通过以下代码验证效果:

for elem in x:
    print(f"uuid: {elem['uuid'].numpy()}, a长度: {len(elem['a'])}, b长度: {len(elem['b'])}")

输出结果符合预期:

uuid: b'abc', a长度: 3, b长度: 3
uuid: b'def', a长度: 3, b长度: 3
uuid: b'ghi', a长度: 3, b长度: 3
uuid: b'pqr', a长度: 4, b长度: 4

方案2:生成器构建数据集

如果后续需要扩展动态加载逻辑、或数据结构更复杂,可使用灵活性更高的生成器方案:

import tensorflow as tf
import numpy as np

def gen():
    t_dic = {"uuid": np.array(["abc", "def", "ghi", "pqr"]),
             "a": [np.array([1, 2, 3]), 
                   np.array([6, 2, 3]), 
                   np.array([6, 8, 1]), 
                   np.array([6, 2, 3, 10])],
             "b": [np.array(["a", "f", "f"]), 
                   np.array(["aa", "ff", "fs"]), 
                   np.array(["aa", "ff", "fs"]), 
                   np.array(["aa", "ff", "fs", "ss"])]}
    for i in range(4):
        yield {
            "uuid": t_dic["uuid"][i],
            "a": t_dic["a"][i],
            "b": t_dic["b"][i]
        }

# 定义输出格式签名
output_signature = {
    "uuid": tf.TensorSpec(shape=(), dtype=tf.string),
    "a": tf.RaggedTensorSpec(shape=[None], dtype=tf.int32),
    "b": tf.RaggedTensorSpec(shape=[None], dtype=tf.string)
}

x = tf.data.Dataset.from_generator(gen, output_signature=output_signature)

该方案最终生成的数据集和方案1效果完全一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 16:36:00