使用TensorFlow Gather处理3D张量的问题求助
解决TensorFlow中按batch维度索引取值的问题
你需要实现对每个batch样本,用3D张量V中的索引提取对应batch行W里的元素,最终得到形状为(P,Q,R)的张量Z。问题出在你没有正确设置batch_dims参数,导致索引逻辑不符合预期。
正确实现代码
import tensorflow as tf # 初始化示例张量 V = tf.random.uniform((2,3,4), minval=0, maxval=2, dtype=tf.int32) W = tf.random.uniform((2,20), minval=0, maxval=4, dtype=tf.int32) # 设置batch_dims=1,让batch维度对应对齐 Z = tf.gather(params=W, indices=V, axis=1, batch_dims=1) print(Z.shape) # 输出 (2, 3, 4),符合预期
为什么原代码不符合要求?
你之前设置batch_dims=0,意味着不考虑batch维度的对应关系,TensorFlow会把整个W当成一个整体进行索引,导致输出多了一个维度(变成4D)。而设置batch_dims=1时,TensorFlow会将W的第一维(batch维度,对应P)和V的第一维对齐,对每个batch行,用V中对应的索引提取W该行的元素,最终输出形状正好是V的形状(P,Q,R)。
另一种实现方式(使用gather_nd)
如果想用gather_nd实现,可以构造包含batch索引和元素索引的二维索引张量:
P, Q, R = V.shape # 生成每个位置对应的batch索引 batch_idx = tf.tile(tf.expand_dims(tf.range(P), axis=(1,2)), (1, Q, R)) # 组合成[batch索引, 元素索引]的结构 indices = tf.stack([batch_idx, V], axis=-1) Z = tf.gather_nd(W, indices) print(Z.shape) # 同样输出 (2, 3, 4)
内容的提问来源于stack exchange,提问作者user3433489
相关产品推荐
相关产品推荐

