PyTorch中如何获取2D张量每行指定最值的索引(向量化实现)
向量化实现按指定规则获取张量每行的最值索引
问题描述
给定形状为(batch_size, N)的正整数张量A(0是张量中的最小值),以及长度为batch_size的列表k:
- 当
k[i] = 1时,取第i行最大值对应的第一个出现的索引 - 当
k[i] = 0时,取第i行非零元素的最小值对应的第一个出现的索引
要求全程使用向量化计算实现,避免逐行循环。
示例
输入张量A:
import torch A = torch.tensor([[4, 3, 1, 4, 2], [0, 0, 2, 3, 4], [4, 4, 3, 0, 3]])
输入k = [1, 0, 0],输出索引为[0, 2, 2],对应值为[4, 2, 3]。
解决方案(向量化实现)
核心思路是先批量计算两种规则下的索引,再根据k的值批量选择结果,全程用PyTorch的张量维度操作完成:
import torch def get_target_indices(A, k): # 将k转为bool型张量并增加维度,方便后续批量选择 k_tensor = torch.tensor(k, dtype=torch.bool).unsqueeze(1) # 批量计算每行最大值的第一个索引 max_indices = A.argmax(dim=1, keepdim=True) # 批量计算每行非零最小值的第一个索引:先把0替换为极大值,再取argmin max_val = A.max() + 1 non_zero_A = torch.where(A == 0, torch.tensor(max_val, dtype=A.dtype), A) min_non_zero_indices = non_zero_A.argmin(dim=1, keepdim=True) # 根据k_tensor批量选择对应索引 output_indices = torch.where(k_tensor, max_indices, min_non_zero_indices).squeeze(1) # 可选:批量获取对应索引的数值 output_values = torch.gather(A, 1, output_indices.unsqueeze(1)).squeeze(1) # 转为列表返回(也可直接返回张量) return output_indices.numpy().tolist(), output_values.numpy().tolist() # 测试示例 A = torch.tensor([[4, 3, 1, 4, 2], [0, 0, 2, 3, 4], [4, 4, 3, 0, 3]]) k = [1, 0, 0] indices, values = get_target_indices(A, k) print(f"output = {indices}") print(f"对应值为 {values}")
关键向量化操作说明
argmax(dim=1)/argmin(dim=1):直接对整个张量按行批量计算最值索引,无需循环torch.where:批量替换0为极大值(排除0对非零最小值的影响),同时批量选择对应规则的索引torch.gather:批量从每行中取出对应索引的数值,实现向量化取值
内容的提问来源于stack exchange,提问作者jupyter
相关产品推荐
相关产品推荐

