如何让TensorFlow DenseHashTable支持多维字符串键查找?
解决DenseHashTable多维字符串键的查找问题
要让tf.lookup.experimental.DenseHashTable支持多维字符串键的查找,实现类似tf.nn.embedding_lookup的效果,核心思路是先将多维键张量展平为一维,完成查找后再恢复原始形状,具体步骤如下:
1. 展平多维键张量
把输入的多维键张量(比如形状(batch_size, 2))通过tf.reshape展平成一维张量,满足DenseHashTable.lookup对键维度的要求。
2. 执行查找操作
用展平后的一维键张量调用table.lookup,得到一维的嵌入结果张量。
3. 恢复原始形状
根据原多维键张量的形状,结合嵌入向量的维度,把查找结果恢复成目标形状(比如原键形状是(batch_size, 2),嵌入维度是4,结果形状应为(batch_size, 2, 4))。
完整代码示例
import tensorflow as tf # 初始化哈希表 keys = ["Fritz", "Franz", "Fred"] values = [[1, 2, 3, -1], [4, 5, -1, -1], [6, 7, 8, 9]] table = tf.lookup.experimental.DenseHashTable( key_dtype=tf.string, value_dtype=tf.float32, empty_key="0", deleted_key="-1", default_value=[-1,-1,-1,-1] ) table.insert(keys, values) # 多维键输入 multi_dimensional_keys = [['Franz', 'Emil'], ['Emil', 'Fred']] key_tensor = tf.convert_to_tensor(multi_dimensional_keys) # 步骤1:展平键张量 flattened_keys = tf.reshape(key_tensor, [-1]) # 步骤2:执行查找 flattened_embeddings = table.lookup(flattened_keys) # 步骤3:恢复原始形状 original_shape = tf.shape(key_tensor) embedding_dim = tf.shape(flattened_embeddings)[1] final_embeddings = tf.reshape(flattened_embeddings, tf.concat([original_shape, [embedding_dim]], axis=0)) # 查看结果 print(final_embeddings.shape) # 输出 (2, 2, 4) print(final_embeddings.numpy())
说明
- 该方法兼容任意维度的键张量,二维、三维等结构都可通过展平-查找-恢复形状的流程处理。
- 哈希表的
default_value会自动应用于不存在的键(比如示例中的"Emil"),和tf.nn.embedding_lookup处理未登录词的逻辑一致。
内容的提问来源于stack exchange,提问作者Eugene Deng
相关产品推荐
相关产品推荐

