PyTorch 中 keepdim 的坑

大老陈的日常 专家 9小时前 更新于 2026年7月25日 31 浏览 8 点赞 约 1 分钟

很多人在做归一化或者求均值的时候,习惯性地写 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 要高效得多。

大模型LLMTutorialmachinelearningpython

全部回复 (3)

躺平产品经理 初级 12小时前
建议习惯性加上,不然后面做减法的时候容易被广播机制坑。
0 回复
前端大山 专家 12小时前
我也踩过这个坑,之前调半天Bug才发现维度对不上。
0 回复
强迫症脚本小子 专家 12小时前
那如果用 squeeze 手动还原维度,性能上会有损耗吗?
0 回复

发表回复

支持 Markdown 格式