如何正确重标记TensorFlow数据集,解决Tensor不可哈希报错问题
报错原因
你触发的TypeError: Tensor is unhashable报错核心原因是:tf.py_function传入自定义函数的输入是Tensor类型,而你直接将Tensor作为Python字典的键做查询,Python字典要求键必须是可哈希的原生数据类型,Tensor不满足该要求。
解决方案
这里提供两种符合你可读性要求的实现方案:
方案1:最小改动适配原有逻辑
仅调整自定义函数的输入处理逻辑,完全保留你原有的字典映射写法,改动量最小:
# 完整可运行代码 import matplotlib.pyplot as plt import numpy as np import tensorflow as tf import tensorflow_datasets as tfds builder = tfds.image.Cifar10() builder.download_and_prepare() ds_train: tf.data.Dataset = builder.as_dataset(split='train') def relabel_map(l): # 先将输入Tensor转为Python原生int类型再查字典 l_py = l.numpy().item() return {0: 0, 1: 6, 2: 1, 3: 2, 4: 3, 5: 4, 6: 9, 7: 5, 8: 7, 9: 8}[l_py] def process_example(example): img = example['image'] label = tf.squeeze(tf.py_function(relabel_map, [example['label']], tf.int64)) # 手动设置标签形状,避免后续数据流水线操作报错 label.set_shape(()) return img, label ds_train = ds_train.map(process_example) for ex in ds_train.take(1): plt.imshow(np.array(ex[0], dtype=np.uint8)) plt.show() print(ex[1])
方案2:TensorFlow原生静态哈希表实现(更推荐)
完全在TensorFlow计算图内执行,性能远高于tf.py_function,适合大数据量训练场景,映射规则直观易读:
import matplotlib.pyplot as plt import numpy as np import tensorflow as tf import tensorflow_datasets as tfds builder = tfds.image.Cifar10() builder.download_and_prepare() ds_train: tf.data.Dataset = builder.as_dataset(split='train') # 定义重映射规则,键为原始标签,值为目标标签 original_labels = tf.constant([0,1,2,3,4,5,6,7,8,9], dtype=tf.int64) target_labels = tf.constant([0,6,1,2,3,4,9,5,7,8], dtype=tf.int64) # 初始化静态哈希表 relabel_table = tf.lookup.StaticHashTable( tf.lookup.KeyValueTensorInitializer(original_labels, target_labels), default_value=-1 ) def process_example(example): return example['image'], relabel_table.lookup(example['label']) ds_train = ds_train.map(process_example) for ex in ds_train.take(1): plt.imshow(np.array(ex[0], dtype=np.uint8)) plt.show() print(ex[1])
内容的提问来源于stack exchange,提问作者leevii
相关产品推荐
相关产品推荐

