如何对Tensor逐行应用tf.gather操作?
逐行对张量执行tf.gather的实现方案
针对你需要对张量每行单独执行tf.gather的需求,这里提供两种高效的实现方式,替代繁琐的tf.scan方案:
方法一:使用tf.gather的batch_dims参数(推荐,TensorFlow 2.0+)
利用tf.gather的batch_dims参数,可以直接指定将前N个维度作为批量维度,自动对每个批量(即每行)执行索引收集操作:
import tensorflow as tf A = tf.constant([[2., 5., 12., 9., 0., 0., 3.], [0., 12., 2., 0., 0., 0., 5.], [0., 0., 10., 0., 4., 4., 3.]], dtype=tf.float32) idxs = tf.constant([[0, 1, 3, 6, 0, 0, 0], [1, 1, 2, 6, 6, 6, 6], [2, 2, 4, 4, 6, 6, 6]], dtype=tf.int32) # 指定batch_dims=1,将每行视为独立批量,用对应idxs行收集元素 output = tf.gather(A, idxs, batch_dims=1) print(output.numpy())
运行后输出:
[[2. 5. 9. 3. 2. 2. 2.] [12. 12. 2. 5. 5. 5. 5.] [10. 10. 4. 4. 3. 3. 3.]]
方法二:手动构造二维索引(兼容旧版本TensorFlow)
如果使用TensorFlow 1.x或不想依赖batch_dims,可以构造每个元素的(行索引, 列索引)对,再用tf.gather_nd收集:
import tensorflow as tf A = tf.constant([[2., 5., 12., 9., 0., 0., 3.], [0., 12., 2., 0., 0., 0., 5.], [0., 0., 10., 0., 4., 4., 3.]], dtype=tf.float32) idxs = tf.constant([[0, 1, 3, 6, 0, 0, 0], [1, 1, 2, 6, 6, 6, 6], [2, 2, 4, 4, 6, 6, 6]], dtype=tf.int32) # 生成与idxs同形状的行索引矩阵 row_indices = tf.tile(tf.range(tf.shape(A)[0])[:, tf.newaxis], [1, tf.shape(idxs)[1]]) # 拼接成(3,7,2)的索引张量,每个元素对应(行,列)位置 gather_indices = tf.stack([row_indices, idxs], axis=-1) # 按索引收集元素 output = tf.gather_nd(A, gather_indices) print(output.numpy())
关于tf.scan的问题
tf.scan是序列式处理操作,虽然能实现需求,但需要手动调整维度且效率低于向量化解法,示例写法如下(不推荐):
import tensorflow as tf A = tf.constant([[2., 5., 12., 9., 0., 0., 3.], [0., 12., 2., 0., 0., 0., 5.], [0., 0., 10., 0., 4., 4., 3.]], dtype=tf.float32) idxs = tf.constant([[0, 1, 3, 6, 0, 0, 0], [1, 1, 2, 6, 6, 6, 6], [2, 2, 4, 4, 6, 6, 6]], dtype=tf.int32) def gather_row(args): row, indices = args return tf.gather(row, indices) # 转置idxs适配tf.scan的处理逻辑,最后再转置回原形状 output = tf.scan(gather_row, (A, tf.transpose(idxs)), axis=0) output = tf.transpose(output) print(output.numpy())
内容的提问来源于stack exchange,提问作者Thomas Ghorbanian
相关产品推荐
相关产品推荐

