You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在对数空间中实现含指数求和的矩阵乘法以避免溢出问题

在对数空间中实现含指数求和的矩阵乘法以避免溢出问题

嘿,我来帮你解决这个溢出难题!你遇到的核心痛点就是直接计算exp(A)和exp(B)时容易触发数值溢出,而且你已经掌握了无sum(dim=1)场景下的对数空间矩阵乘法解法,但不知道怎么扩展到带求和的情况。其实我们可以把整个运算全流程都放在对数空间里完成,完全避开直接计算指数,具体步骤如下:

核心思路拆解

你的目标表达式是:

out = torch.log(torch.exp(A).sum(dim=1) @ torch.exp(B).sum(dim=1))

我们可以把它拆成三个关键部分,全部用对数空间操作替代:

  1. 把torch.exp(A).sum(dim=1)转换成对数空间表示:这其实就是PyTorch内置torch.logsumexp函数的本职工作——log(exp(X).sum(dim=d))等价于torch.logsumexp(X, dim=d),完美绕开直接计算exp。
  2. 用同样的方式处理torch.exp(B).sum(dim=1)。
  3. 最关键的一步:计算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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 14:24:34