如何用索引数组从张量型数组数组中选取对应元素
嘿,这个需求其实挺常见的,我来给你几种实用的实现方式,你可以根据自己的场景来选:
方法1:纯Python列表推导式(小数据量首选)
如果你的数据就是普通的Python嵌套列表,用列表推导式最直观,代码也简洁:
# 定义原始数组和索引数组 arr = [[1,1],[3,1],[4,5],[7,2],[6,0]] indices = [0,0,1,1,0] # 用zip配对每行和对应的索引,取出元素后包裹成单元素列表 result = [[row[idx]] for row, idx in zip(arr, indices)] print(result) # 输出: [[1], [3], [5], [2], [6]]
这种方式逻辑清晰,适合数据量不大的场景,不用额外依赖库。
方法2:NumPy实现(张量/大数据量高效方案)
如果你的“张量型数组”是指NumPy数组(机器学习里常用的张量形式),那用NumPy的索引机制会更高效,尤其是数据量大的时候:
import numpy as np # 转成NumPy数组 arr_np = np.array([[1,1],[3,1],[4,5],[7,2],[6,0]]) indices_np = np.array([0,0,1,1,0]) # 利用高级索引取出对应元素,再reshape成单元素数组的形式 selected = arr_np[np.arange(len(arr_np)), indices_np] result = selected.reshape(-1, 1).tolist() # 转成Python列表格式 print(result) # 输出: [[1], [3], [5], [2], [6]]
这里np.arange(len(arr_np))生成每行的索引,和indices_np配对,就能精准定位到每个行里要取的元素,reshape(-1,1)把一维数组转成列向量,最后用tolist()转回Python列表就和期望输出一致了。
额外补充:PyTorch张量的情况
如果你的张量是PyTorch的Tensor,实现逻辑和NumPy几乎一样:
import torch arr_tensor = torch.tensor([[1,1],[3,1],[4,5],[7,2],[6,0]]) indices_tensor = torch.tensor([0,0,1,1,0]) selected = arr_tensor[torch.arange(arr_tensor.size(0)), indices_tensor] result = selected.reshape(-1, 1).tolist() print(result) # 输出同样符合要求
内容的提问来源于stack exchange,提问作者Arijit Gupta
相关产品推荐
相关产品推荐

