
本文介绍一种避免数值溢出的高效方法,用于计算形如 log(sum(exp(A), dim=1) @ sum(exp(B), dim=1)) 的表达式,通过重写为对数空间中的 logsumexp 运算实现全程数值稳定。
本文介绍一种避免数值溢出的高效方法,用于计算形如 `log(sum(exp(a), dim=1) @ sum(exp(b), dim=1))` 的表达式,通过重写为对数空间中的 `logsumexp` 运算实现全程数值稳定。
在深度学习与概率建模中,我们常需对高维张量执行指数运算后求和再相乘(例如在隐变量模型或注意力机制变体中),但直接计算 torch.exp(A) 极易因中间值过大导致 inf 溢出,破坏梯度传播与数值一致性。原始表达式:
out = torch.log(torch.exp(A).sum(dim=1) @ torch.exp(B).sum(dim=1))
虽语义清晰,却无法规避 exp 阶段的数值不稳定性。关键洞察在于:矩阵乘法本质是双重复合求和——即 (X @ Y)[i,k] = Σ_j X[i,j] * Y[j,k]。当 X = sum(exp(A), dim=1)、Y = sum(exp(B), dim=1) 时,整个运算可完全等价地在对数空间展开。
设 log_S_A = log(Σₗ exp(A[l,i,j]))(即 torch.logsumexp(A, dim=1)),同理得 log_S_B。则目标结果第 (i,k) 项为:
log( Σ_j exp(log_S_A[i,j]) * exp(log_S_B[j,k]) ) = log( Σ_j exp(log_S_A[i,j] + log_S_B[j,k]) ) = logsumexp_j (log_S_A[i,j] + log_S_B[j,k])
这正是 torch.logsumexp 的标准适用场景。因此,只需将 log_S_A 在 j 维(即第二个维度)扩展为 (bs, m, 1, m),log_S_B 在 j 维(即第一个维度)扩展为 (bs, 1, m, m),二者广播相加得到 (bs, m, m, m) 的中间张量,最后沿 j 维(dim=2)做 logsumexp 即可。
以下是完整、可直接复用的稳定实现:
import torch
def stable_log_sum_exp_matmul(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
稳定计算: log( sum(exp(A), dim=1) @ sum(exp(B), dim=1) )
输入:
A, B: shape (bs, n, m, m)
输出:
out: shape (bs, m, m),满足 out[b,i,k] ≈
log( Σ_j [Σ_l exp(A[b,l,i,j])] * [Σ_l' exp(B[b,l',j,k])] )
"""
# Step 1: 对 batch 内第1维(n)求 log-sum-exp → 形状变为 (bs, m, m)
log_S_A = torch.logsumexp(A, dim=1) # (bs, m, m)
log_S_B = torch.logsumexp(B, dim=1) # (bs, m, m)
# Step 2: 广播相加:log_S_A[i,j] + log_S_B[j,k] → 需对齐 j 维
# 将 log_S_A 扩展为 (bs, m, 1, m),log_S_B 扩展为 (bs, 1, m, m)
combined = log_S_A.unsqueeze(2) + log_S_B.unsqueeze(1) # (bs, m, m, m)
# Step 3: 沿 j 维(即 dim=2)做 logsumexp → 得到 (bs, m, m)
out = torch.logsumexp(combined, dim=2)
return out✅ 优势总结:
- 全程不调用
torch.exp或torch.log于大数值,彻底规避inf/nan; - 时间复杂度与原式一致(O(bs·m³)),仅引入少量广播开销;
- 支持自动微分,梯度计算同样稳定;
- 可无缝集成至 PyTorch 训练流程(如自定义 loss 或 attention kernel)。
⚠️ 注意事项:
- 输入
A和B应为float32或float64张量;若含nan/inf,logsumexp会传播异常,建议前置检查; - 当
n=1时,该函数退化为稳定版log(exp(A) @ exp(B)),与 Stack Overflow 中的经典解法一致; - 若内存受限(如
m > 256),可考虑分块logsumexp或使用torch.einsum替代广播(但通常无必要)。
此方法体现了“将数值敏感操作整体重参数化至对数域”的核心思想,是处理指数族分布、softmax 变体及对数空间线性代数任务的标准范式。

















