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

Numpy指数溢出与标量除法无效值问题求解(不损失精度)

解决Softmax计算中的指数溢出问题

问题分析

你的prob函数本质是计算类Softmax的概率分布,但直接计算np.exp()时,当phi*(a3-a2)取值较大(比如phi=1、a3=30、a2=12时,该项为18),若优化过程中phi临时超出0-1范围,或处理批量数据时累计计算,就可能触发指数溢出,进而导致除法运算出现无效值。

解决方案:Log-Sum-Exp技巧

避免指数溢出的标准方法是使用**Log-Sum-Exp(LSE)**技巧,核心思路是通过减去所有指数项的最大值,将所有指数参数压缩到非正区间,保证np.exp()的结果不会超过1,彻底避免溢出。推导逻辑如下:

对于分母 de = exp(x1) + exp(x2) + exp(x3)(其中x1=phi*(a1-a2)、x2=0、x3=phi*(a3-a2)),可改写为:

de = exp(max_x) * [exp(x1-max_x) + exp(x2-max_x) + exp(x3-max_x)]

对应的分子exp(xi)为exp(max_x)*exp(xi-max_x),两者相除时exp(max_x)会抵消,最终概率为:

pi = exp(xi - max_x) / [exp(x1-max_x) + exp(x2-max_x) + exp(x3-max_x)]

该计算方式与原逻辑完全等价,且无溢出风险。

修改后的代码

import numpy as np

def attraction(b1, b2, b3):
    a1 = 10 * b1 + 20 * b2 + 6 * b3
    a2 = 12 * b1 + 18 * b2 + 10 * b3
    a3 = 0 * b1 + 10 * b2 + 30 * b3
    return a1, a2, a3

def prob(a1, a2, a3, phi):
    # 定义三个logit项
    x1 = phi * (a1 - a2)
    x2 = 0.0
    x3 = phi * (a3 - a2)
    
    # 获取最大logit,压缩指数范围避免溢出
    max_x = np.max([x1, x2, x3])
    
    # 计算归一化后的exp项
    exp_x1 = np.exp(x1 - max_x)
    exp_x2 = np.exp(x2 - max_x)
    exp_x3 = np.exp(x3 - max_x)
    
    de = exp_x1 + exp_x2 + exp_x3
    p1 = exp_x1 / de
    p2 = exp_x2 / de
    p3 = exp_x3 / de
    
    return p1, p2, p3

额外说明

  • 该方法无需依赖float128类型,完全兼容M1 Mac的numpy环境,且不会损失计算精度;
  • 即使phi临时超出0-1的预期范围,也能稳定计算,避免溢出警告;
  • 若处理批量数据(如a1/a2/a3为数组),numpy的广播机制会自动适配,无需额外修改。

内容的提问来源于stack exchange,提问作者jasmine

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 14:20:25