TensorFlow中tf.nn.embedding_lookup无法拦截无效索引问题求助
为什么
tf.where没能阻止无效索引传入tf.nn.embedding_lookup? 这是个很典型的TensorFlow图执行模式的坑!问题出在TensorFlow的图是预先构建并全部分支执行的,tf.where并不会惰性跳过不满足条件的分支——哪怕你写了“当k=-1时返回零张量”,tf.nn.embedding_lookup(phi, k)这个操作还是会在图执行阶段被完整计算,当k里出现-1时,自然就触发了索引越界的错误。
解决思路:先处理无效索引,再做嵌入查找
我们可以分三步绕开这个问题:
- 把所有-1的索引替换成一个合法的索引(比如0,后面再把它的结果清零)
- 执行嵌入查找(此时所有索引都是合法的,不会报错)
- 用掩码把原本是-1的位置的结果置为零
对应的代码实现:
# 第一步:将k中的-1替换为合法索引(这里选0,任意合法索引都可以) k_safe = tf.where(tf.equal(k, -1), tf.zeros_like(k), k) # 第二步:执行嵌入查找,此时所有索引都在有效范围内 lookup_result = tf.nn.embedding_lookup(phi, k_safe, name="k_cluster_look") # 第三步:生成掩码,将k=-1的位置标记为0,其余为1(扩展维度匹配嵌入的Y维度) mask = tf.cast(tf.not_equal(k, -1), tf.float32)[:, tf.newaxis] # 掩码与查找结果相乘,-1对应的位置就变成全零张量 kappa = lookup_result * mask
验证效果
当你执行:
x = session.run(kappa, feed_dict={k:[-1]})
此时k_safe会变成[0],lookup_result取出phi[0],然后mask是[[0.]],相乘后得到的就是和tf.zeros([1,self.Y])完全一致的全零张量,完美符合你的需求;如果传入批量索引比如k=[3, -1, 7],对应的结果会正确保留phi[3]、全零、phi[7]。
补充说明
如果是TensorFlow 2.x的即时执行(Eager Execution)模式,虽然tf.where会根据条件选择执行分支,但tf.nn.embedding_lookup碰到-1还是会直接报错,所以上面的方案同样适用——先把无效索引转成合法值,再用掩码修正结果,是通用的解决办法。
内容的提问来源于stack exchange,提问作者KannarKK
相关产品推荐
相关产品推荐

