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

tf.data.Dataset.filter无法访问字典,如何修复张量不可哈希问题?

修复TensorFlow数据集过滤时标签张量无法作为字典键的问题

问题场景

为平衡图像数据集,先统计各标签的图像数量并找到单标签最小数量,编写过滤函数移除多余图像时,出现以下错误:

TypeError: Tensor is unhashable. Instead, use tensor.ref() as the key.

错误出现在尝试用张量类型的标签作为Python字典的键时。

核心原因

TensorFlow在图执行模式下,传入过滤函数的label是张量对象,而非原生Python类型(如整数)。Python字典要求键是可哈希的原生类型,因此直接用张量作为键会触发错误。

解决方案

方案1:转换张量为原生Python类型(简单直接)

将张量标签转换为Python原生值后再操作字典,同时用tf.py_function包裹过滤逻辑,确保Python代码能在图模式中执行:

import tensorflow as tf

# 假设已统计好各标签数量,label_count_min是单标签最小数量
label_counts = {0: 200, 1: 150, 2: 180}
label_count_min = 150

def filter_fn(image, label):
    # 将张量标签转为Python整数
    label_py = label.numpy()
    count = label_counts[label_py]
    keep = count <= label_count_min
    
    if keep:
        label_counts[label_py] -= 1
    return keep

# 用tf.py_function包裹过滤函数,指定输出类型为布尔值
filtered_dataset = dataset.filter(
    lambda img, lbl: tf.py_function(func=filter_fn, inp=[img, lbl], Tout=tf.bool)
)

方案2:使用TensorFlow变量存储计数(更贴合图模式)

避免Python字典的线程安全问题,改用TensorFlow变量存储各标签的计数,全程用TensorFlow操作:

import tensorflow as tf

label_counts = {0: 200, 1: 150, 2: 180}
label_count_min = 150

# 将标签键和计数转为TensorFlow格式
label_keys = tf.constant(list(label_counts.keys()), dtype=tf.int32)
count_vars = tf.Variable(list(label_counts.values()), dtype=tf.int32)

def filter_fn(image, label):
    # 找到当前标签在keys中的索引
    idx = tf.where(tf.equal(label_keys, label))[0][0]
    current_count = count_vars[idx]
    keep = current_count <= label_count_min
    
    # 原子性减少计数(线程安全)
    update = tf.one_hot(idx, depth=len(label_keys), dtype=tf.int32)
    count_vars.assign_sub(update)
    
    return keep

filtered_dataset = dataset.filter(filter_fn)

注意事项

  • 方案1中,label_counts需是全局或闭包内的可变对象,确保过滤过程中计数能被正确更新。
  • 若数据集规模大、多线程处理,方案2的TensorFlow变量方式更安全,避免Python字典的并发修改问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 21:03:37