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
相关产品推荐
相关产品推荐

