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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 16:15:13