如何在非Eager Execution模式下加载TensorFlow数据集
TensorFlow图模式下interleave加载多份已保存数据集报错解决方案
问题描述
需在关闭Eager Execution的图模式场景下加载多份本地保存的TensorFlow数据集,已知条件:
- 已拿到所有数据集的存储路径列表
- 已提前获取每个路径对应数据集的
element_spec - 目标是加载所有数据集后通过
interleave操作合并为统一数据集
故障表现:
- Eager模式下通过for循环遍历路径逐个加载可正常运行
- 用
map/interleave实现图模式加载时触发报错:TypeError: Tensor is unhashable. Instead, use tensor.ref() as the key.
问题复现代码
import tensorflow as tf # 构造测试数据集并保存 data1 = tf.data.Dataset.from_tensor_slices(([3, 4], [0, 1])) print(list(data1.as_numpy_iterator())) data1.save(path='1') data2 = tf.data.Dataset.from_tensor_slices(([5, 6], [2, 3])) print(list(data2.as_numpy_iterator())) data2.save(path='2') # 构造路径数据集 dataset = tf.data.Dataset.from_tensor_slices(['1', '2']) elem_specs = {0: data1.element_spec, 1: data2.element_spec} dataset = dataset.enumerate() # Eager模式下可正常运行 for i, ds in dataset: result = tf.data.Dataset.load(ds, element_spec=elem_specs[i.numpy()]) # 图模式下interleave调用触发报错 result = dataset.interleave(lambda i, item: tf.data.Dataset.load(item, element_spec=elem_specs[i])) # TypeError: Tensor is unhashable. Instead, use tensor.ref() as the key.
报错原因
interleave/map传入的处理函数运行在图模式下,函数内的索引i是未求值的Tensor对象,不是Python原生整数。此时直接用i作为键查询Python原生字典elem_specs,Python侧会尝试对Tensor对象做哈希操作,直接触发类型错误。
可行实现方案
方案1:提前配对路径与spec,用生成器构造数据集(推荐)
不在图运行时做Python字典查询,提前将路径和对应的element_spec绑定,通过生成器逐个加载子数据集后再做interleave合并,兼容性最好:
import tensorflow as tf tf.compat.v1.disable_eager_execution() # 显式关闭Eager验证图模式场景 # 提前按顺序绑定路径和对应的element_spec path_spec_pairs = [ ('1', data1.element_spec), ('2', data2.element_spec) ] # 定义逐个子数据集加载的生成器 def load_sub_ds(): for path, spec in path_spec_pairs: yield tf.data.Dataset.load(path, element_spec=spec) # 构造子数据集集合,interleave合并为统一数据集 merged_dataset = tf.data.Dataset.from_generator( load_sub_ds, output_signature=tf.data.DatasetSpec.from_value(data1) ).interleave(lambda sub_ds: sub_ds, cycle_length=2) # 验证输出 print(list(merged_dataset.as_numpy_iterator())) # 输出: [(3,0), (5,2), (4,1), (6,3)]
方案2:用TensorFlow原生分支接口替代Python字典查询
如果必须在interleave逻辑内根据索引动态选择spec,使用图兼容的tf.switch_case接口实现分支,不要用Python原生字典做查询:
dataset = tf.data.Dataset.from_tensor_slices(['1', '2']).enumerate() def load_branch(idx, path): # 定义分支索引与对应加载逻辑 branch_map = { 0: lambda: tf.data.Dataset.load(path, element_spec=data1.element_spec), 1: lambda: tf.data.Dataset.load(path, element_spec=data2.element_spec) } return tf.switch_case(idx, branch_map) merged_dataset = dataset.interleave(load_branch, cycle_length=2)
注意事项
- 不要在
tf.data的图运行逻辑中,使用未求值的Tensor作为键操作Python原生字典、列表等容器,这类操作在Python侧执行,无法识别Tensor的运行时值 - 所有依赖Tensor具体值的分支、查询逻辑,优先使用TensorFlow原生提供的图兼容算子实现
内容的提问来源于stack exchange,提问作者Jesse Kerr
相关产品推荐
相关产品推荐

