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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 09:35:21