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

是否存在numpy函数可提取矩阵每行指定列元素生成新数组?

Numpy快速实现方案

完全可以通过numpy高级索引实现你要的逻辑,全程没有Python层面的循环,性能远高于列表推导,数据量越大优势越明显。

最常用的arange索引方案

这是最简洁也最高效的实现方式,和你原有逻辑完全等价:

import numpy as np
correct_answers = scores[np.arange(num_train), y]

实现原理

  • np.arange(num_train)会生成长度为num_train的数组,元素为0到num_train-1,刚好对应scores的所有行索引
  • 当传入两个长度相同的一维数组作为索引时,numpy会逐位置配对取数:即依次取(arange[i], y[i])位置的元素,和你写的列表推导逻辑完全一致
  • 整个运算在numpy底层C层面完成,不会有Python循环的额外开销

可选替代方案

如果需要适配更高维度的场景,也可以用更易读的take_along_axis实现:

# 需要先把y转为列向量,指定沿列维度取数,最后打平为一维数组
correct_answers = np.take_along_axis(scores, y.reshape(-1, 1), axis=1).flatten()

效果验证

可以用下面的测试样例确认两种方法结果完全一致:

# 构造测试数据
num_train = 3
columns = 4
scores = np.arange(num_train * columns).reshape(num_train, columns)
y = np.array([1, 3, 0])

# 原有列表推导实现
correct_answers_loop = np.array([scores[i][y[i]] for i in range(num_train)])
# 高级索引实现
correct_answers_fast = scores[np.arange(num_train), y]

print(correct_answers_loop) # 输出 [1 7 8]
print(correct_answers_fast) # 输出 [1 7 8]

内容的提问来源于stack exchange,提问作者Calin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 19:27:04