在对数空间中实现含指数求和的矩阵乘法以避免溢出问题
在对数空间中实现含指数求和的矩阵乘法以避免溢出问题
嘿,我来帮你解决这个溢出难题!你遇到的核心痛点就是直接计算exp(A)和exp(B)时容易触发数值溢出,而且你已经掌握了无sum(dim=1)场景下的对数空间矩阵乘法解法,但不知道怎么扩展到带求和的情况。其实我们可以把整个运算全流程都放在对数空间里完成,完全避开直接计算指数,具体步骤如下:
核心思路拆解
你的目标表达式是:
out = torch.log(torch.exp(A).sum(dim=1) @ torch.exp(B).sum(dim=1))
我们可以把它拆成三个关键部分,全部用对数空间操作替代:
- 把
torch.exp(A).sum(dim=1)转换成对数空间表示:这其实就是PyTorch内置torch.logsumexp函数的本职工作——log(exp(X).sum(dim=d))等价于torch.logsumexp(X, dim=d),完美绕开直接计算exp。 - 用同样的方式处理
torch.exp(B).sum(dim=1)。 - 最关键的一步:计算
log(S_A @ S_B)(其中S_A = exp(A).sum(dim=1),S_B = exp(B).sum(dim=1))。这里的矩阵乘法+对数,本质上是对每一对行和列的元素和做logsumexp。
具体实现代码
直接上可运行的PyTorch代码,全程无溢出风险:
import torch # 假设A和B的形状是(bs, n, m, m) # 第一步:替代直接exp+sum,用logsumexp得到对数空间的求和结果 log_S_A = torch.logsumexp(A, dim=1) # 形状变为(bs, m, m) log_S_B = torch.logsumexp(B, dim=1) # 形状同样是(bs, m, m) # 第二步:用广播+logsumexp实现对数空间的矩阵乘法 # 扩展维度让log_S_A的行和log_S_B的列可以逐元素相加 # log_S_A.unsqueeze(2) → (bs, m, 1, m),log_S_B.unsqueeze(1) → (bs, 1, m, m) log_product = log_S_A.unsqueeze(2) + log_S_B.unsqueeze(1) # 形状为(bs, m, m, m) # 对中间的维度(对应矩阵乘法的求和维度)做logsumexp,得到最终结果 out = torch.logsumexp(log_product, dim=2) # 形状回到(bs, m, m),和原表达式结果一致
为什么这能行?
从数学上验证一下逻辑的正确性:
log_S_A = torch.logsumexp(A, dim=1)严格等价于torch.log(torch.exp(A).sum(dim=1)),而且logsumexp会通过先减去最大值再计算的数值稳定策略,从根源上避免溢出。- 矩阵乘法
S_A @ S_B的每个元素(i,j)是sum_k S_A[i,k] * S_B[k,j],转换成对数空间就是log(sum_k exp(log_S_A[i,k] + log_S_B[k,j]))——这正好是我们通过广播相加后,对k维度(代码中的dim=2)做logsumexp的结果。
测试验证
你可以用小数值的A和B做测试,比如把A/B的元素限制在[-5,5]之间(不会触发溢出),对比原表达式和新代码的结果,应该会在浮点误差范围内完全一致;同时当A/B元素数值很大时,新代码也不会出现溢出报错。
备注:内容来源于stack exchange,提问作者vendrick17
相关产品推荐
相关产品推荐

