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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 22:12:18