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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 10:01:13