PyTorch 中 keepdim 的坑
很多人在做归一化或者求均值的时候,习惯性地写
下一篇
代码接手最怕遇到“消失的救星”,尤其是在项目烂摊子堆成山的时候。 →
x.sum(dim=1),结果要么直接报个 RuntimeError 让你怀疑人生,要么更恶心——代码跑通了,但算出来的数值完全不对。这时候如果有人告诉你加个 keepdim=True 就能解决,你可能觉得是某种“魔法”,但其实这背后是 PyTorch 的广播机制在搞鬼。简单来说,所有的还原操作(sum, mean, max, min)默认都会把指定的维度“压扁”并直接删掉。比如一个 (2, 3) 的张量,你在 dim=1 上求和,结果会变成 (2,)。而如果你设置了 keepdim=True,那个维度会被保留成 1,结果就是 (2, 1)。
看这个实操对比:
import torch
m = torch.tensor([[1., 2., 3.],
[4., 5., 6.]])
# 默认情况:维度消失了
print(m.sum(dim=1).shape) # torch.Size([2])
# 开启 keepdim:维度变成了 1
print(m.sum(dim=1, keepdim=True).shape) # torch.Size([2, 1])真正关键的地方在于接下来的“广播(Broadcasting)”。
当你需要用原张量除以这个和(比如做行归一化)时,PyTorch 会从右向左对齐维度。(2, 3) 对齐 (2, 1) 时,那个 1 会被自动拉伸成 3,计算完美契合。但如果对齐的是 (2,),PyTorch 发现末尾是 3 vs 2,直接崩溃。
# 正确做法:keepdim=True 保证能够广播
normed = m / m.sum(dim=1, keepdim=True)
print(normed.sum(dim=1)) # tensor([1., 1.]) ✅ 每一行和为 1最阴险的是,如果你的张量刚好是方阵(比如 (2, 2)),不加 keepdim 可能也不会报错,但它会按照广播规则在错误的维度上进行计算,导致结果在逻辑上完全错误,而且没有任何提示。
所以,只要涉及到“先降维求值,再回原张量做运算”的场景,闭眼加上 keepdim=True 准没错,这比事后去调试那些莫名其妙的数值 Bug 要高效得多。