如何使用torch.min获取shape=(x,1)张量的最小值及对应索引
torch.min获取最小值及对应索引的正确用法 torch.min的返回规则如下:
- 不传入
dim参数直接调用时,仅返回张量全局最小值的标量结果,不携带索引信息。 - 传入
dim参数指定计算维度时,返回二元元组:第一个元素是指定维度上的最小值张量,第二个元素是最小值在该维度上的位置索引张量。
针对shape为(x, 1)的二维张量场景,沿第0维度(行方向)计算最小值,即可得到和示例完全匹配的输出:
import torch a = torch.tensor([[10], [5], [8], [2], [8]]) # 沿dim=0(行方向)规约计算最小值 min_value, min_index = torch.min(a, dim=0) print(min_value) # 输出: tensor([2]) print(min_index) # 输出: tensor([3])
其他场景适配
如果需要对任意形状的张量取全局最小值和对应的全局一维索引,可以先将张量展平后计算:
# 展平为一维张量后取最小值与索引 min_value, min_index = torch.min(a.flatten(), dim=0) # 若需要保持示例中(1,)的输出形状,追加一维即可 min_value = min_value.unsqueeze(0) min_index = min_index.unsqueeze(0)
注意:如果张量内存在多个相等的最小值,torch.min默认返回索引最小(最先出现)的位置。
内容的提问来源于stack exchange,提问作者aktabit
相关产品推荐
相关产品推荐

