如何基于含不等长列表的字典创建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
相关产品推荐
相关产品推荐

