如何基于元素个数比较两个PyTorch张量的尺寸大小
问题说明
需求为对比两个PyTorch张量(以1维张量为主要使用场景)的尺寸,判断依据为张量包含的元素总数量,最终返回元素总数更大的张量。
问题复现
直接调用torch.maximum()接口无法实现上述需求,该接口的设计目标不是对比张量尺寸,传入空张量时会返回不符合预期的结果,复现代码如下:
>>> import torch >>> tensor1 = torch.empty(0) >>> tensor2 = torch.empty(1) >>> tensor1 tensor([]) >>> tensor2 tensor([5.9555e-34]) >>> torch.maximum(tensor1,tensor2) tensor([])
错误原因
torch.maximum()是逐元素数值比较接口,仅用于对比两个可广播张量对应位置的数值大小,返回对应位置更大值组成的新张量,不支持张量本身尺寸、元素总数这类元属性的对比- 空张量参与广播运算时的规则导致了上述反常输出,本质是接口选型错误,不属于PyTorch框架bug。
正确实现方案
通过PyTorch张量原生的.numel()方法获取张量的元素总个数,对比两个张量的元素个数后返回对应张量即可,实现代码如下:
def get_larger_tensor(t_a, t_b): # 对比两个张量的元素总数,返回元素数更多的张量 # 元素数相等时默认返回第一个传入的张量 return t_a if t_a.numel() >= t_b.numel() else t_b
效果验证
针对上述复现场景测试,返回结果符合预期:
>>> tensor1 = torch.empty(0) >>> tensor2 = torch.empty(1) >>> get_larger_tensor(tensor1, tensor2) tensor([5.9555e-34])
补充说明
.numel()方法对任意维度的PyTorch张量都生效,不局限于1维张量场景- 如果需要对比2个以上的张量,遍历所有张量取
.numel()值最大的对应张量即可 - 所有逐元素运算类接口的处理对象都是张量内存储的数值,不要用这类接口处理张量本身的元属性对比需求。
内容的提问来源于stack exchange,提问作者Shijith
相关产品推荐
相关产品推荐

