TensorFlow哈希表处理数组类型键值报错:形状不兼容问题求助
解决TensorFlow HashTable使用数组键时的形状不兼容错误
这个问题我之前也碰到过,核心原因是输入张量的形状和哈希表中键的形状不匹配,导致lookup操作对默认值的形状预期出错了。
问题分析
你的代码里:
keys是形状为(3,2)的张量,每个键是长度为2的一维数组;- 但传入
lookup的input_tensor是形状为(2,)的一维张量,TensorFlow会把它解析成两个独立的标量键(1和1)去查询,而不是一个完整的长度为2的数组键; - 此时哈希表期望默认值是标量形状
(),但你提供的default_value是形状(2,)的张量,两者形状不兼容,就触发了ValueError: Shapes (2,) and () are not compatible。
解决方案
只需要把input_tensor转换成和单个键同形状的批量张量——给它增加一个维度,让它的形状和keys中的每个元素一致(即(1,2)),这样lookup就会把整个数组当作一个完整的键去匹配。
修改后的完整代码:
import numpy as np import tensorflow as tf # 给input_tensor增加一个维度,将形状从(2,)变为(1,2),匹配单个键的形状 input_tensor = tf.expand_dims(tf.constant([1, 1], dtype=tf.int64), 0) keys = tf.constant(np.array([[1, 1],[2, 2],[3, 3]]), dtype=tf.int64) values = tf.constant(np.array([[4, 1],[5, 1],[6, 1]]), dtype=tf.int64) default_value = tf.constant(np.array([1, 1]), dtype=tf.int64) table = tf.contrib.lookup.HashTable( tf.contrib.lookup.KeyValueTensorInitializer(keys, values), default_value ) out = table.lookup(input_tensor) with tf.Session() as sess: table.init.run() print(out.eval()) # 输出结果:[[4 1]],符合预期
扩展说明
如果需要批量查询多个数组键(比如要查询[[1,1], [2,2]]),只需要保持input_tensor的形状为(2,2)即可,lookup会自动批量处理每个数组键,返回对应的values或default_value。
内容的提问来源于stack exchange,提问作者Mihkel L.
相关产品推荐
相关产品推荐

