TensorFlow/Numpy中如何基于列索引矩阵按行收集矩阵元素
实现方案
这个按行匹配列索引取值的需求,不需要写循环遍历,TensorFlow 和 NumPy 都有原生支持的实现方式,执行效率比逐行拼接高很多。
TensorFlow 实现
最简洁的写法是直接用 tf.gather,指定 batch_dims=1 即可,逻辑和你提到的 PyTorch torch.gather(dim=1) 完全一致,开箱即用:
import tensorflow as tf # 沿用你给出的测试样例 test1 = tf.constant([[1., 1., 2.], [4., 5., 0.]], dtype=tf.float32) test_ind = tf.constant([[0,1,0,0,1], [0,1,1,1,0]], dtype=tf.int64) result = tf.gather(params=test1, indices=test_ind, axis=1, batch_dims=1)
运行得到的输出和你给出的预期结果完全匹配:
<tf.Tensor: shape=(2, 5), dtype=float32, numpy= array([[1., 1., 1., 1., 1.], [4., 5., 5., 5., 4.]], dtype=float32)>
如果需要更灵活的索引控制,也可以用 tf.gather_nd 手动构造完整坐标索引,逻辑是给每一个列索引配上对应的行号,组成(行,列)二维坐标后直接取值:
def gather_matrix_indices_tf(input_arr, index_arr): # 生成与索引张量形状一致的行号矩阵 row_ind = tf.range(tf.shape(index_arr)[0], dtype=index_arr.dtype)[:, None] row_ind = tf.broadcast_to(row_ind, tf.shape(index_arr)) # 拼接为完整的N维索引 full_ind = tf.stack([row_ind, index_arr], axis=-1) return tf.gather_nd(input_arr, full_ind)
NumPy 实现
NumPy 可以直接用原生高级索引完成相同逻辑,不需要调用额外函数:
import numpy as np def gather_matrix_indices_np(input_arr, index_arr): row_ind = np.arange(index_arr.shape[0])[:, None] return input_arr[row_ind, index_arr]
内容的提问来源于stack exchange,提问作者AndrewJaeyoung
相关产品推荐
相关产品推荐

