PyTorch torch.max文档未明确双张量传入逻辑及相关报错疑问
关于torch.max用法的两个疑问解答
问题背景
我在学习PyTorch与深度学习时,发现torch.max的官方文档没明确说明可以传入两个张量做比较,但实际代码里有这种用法。最初我用h1 = torch.max(h1, 0)实现ReLU逻辑,触发了TypeError: unsupported operand type(s) for @: 'torch.return_types.max' and 'Tensor'报错,改成传入torch.tensor(0)后问题解决,有两个疑问需要解答。
1. 为何该修改能修复报错?
torch.max有两种核心调用逻辑:
- 按维度取最大值:调用格式是
torch.max(input, dim),返回的是一个包含(最大值张量, 索引张量)的元组(类型为torch.return_types.max)。当你传入h1和整数0时,PyTorch会把0识别成dim(维度参数),因此返回的是这个元组。之后h1被赋值为该元组,再执行矩阵乘法@操作时,自然会报错——元组无法和张量进行乘法运算。 - 逐元素取两个输入的最大值:调用格式是
torch.max(input, other),要求两个输入是可广播的张量(或能被正确识别为张量的对象)。当你传入torch.tensor(0)时,PyTorch会判定这是两个张量的逐元素比较场景,返回的是和h1同形状的张量,后续的@操作就能正常执行了。
2. 根据现有文档,torch.tensor(0)属于何种参数?
在PyTorch官方文档中,torch.tensor(0)对应的是torch.max(input, other, *, out=None)接口里的other参数。这个参数需要是一个能和input进行广播运算的张量(或符合广播规则的标量张量),作用是与input逐元素比较,保留每个位置上的较大值。
内容的提问来源于stack exchange,提问作者heretoinfinity
相关产品推荐
相关产品推荐

