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

如何借助Numba进一步加速多维logsumexp与softmax计算

优化Numba实现的Softmax/LogSumExp计算

现有代码的问题分析

  • 过细的循环嵌套:最内层循环逐个处理3个元素,没有利用Numba对数组批量操作的优化能力,反而增加了循环调度开销。
  • Numpy函数冗余调用:循环内使用np.max、np.log等Numpy函数,这类函数在Numba循环内会产生额外调用开销,且阻碍fastmath优化生效——Numba的fastmath对纯Python内置函数或原生操作的优化更显著。
  • 内存访问效率低:循环内频繁做单个元素赋值,加上数组转置的额外开销,导致缓存命中率下降。

优化方案与代码实现

核心优化点

  1. 批量处理固定长度的内层维度(3个元素),避免逐个元素赋值
  2. 替换Numpy函数为math模块的原生函数,让fastmath优化生效
  3. 缓存子数组减少索引计算开销,提升缓存命中率
  4. 可选:利用Numba并行化外层循环,压榨多核CPU性能

基础优化版本

import numba
import numpy as np
import math

def get_p_4d(a, lamda):
    m = a * lamda[:, None][:,None].transpose(0,3,1,2)
    c = np.max(m, axis=2)[:,None].transpose(0,2,1,3)
    aa = np.exp(m - c)
    logsumexp = c + np.log(aa.sum(axis=2)[:,None].transpose(0,2,1,3))
    p = np.exp(m - logsumexp)
    return p

@numba.njit(fastmath=True)
def get_p_4d_nb_opt(a, lamda):
    num_code, num_draw, _, _ = a.shape
    # 提前完成数组转置,避免循环内重复操作
    a_trans = a.transpose(0, 1, 3, 2)
    p = np.empty((num_code, num_draw, 3, 3), dtype=a.dtype)
    
    for i in range(num_code):
        # 缓存子数组,减少索引层级
        lamda_i = lamda[i]
        a_i = a_trans[i]
        p_i = p[i]
        for j in range(num_draw):
            this_lamda = lamda_i[j]
            a_ij = a_i[j]
            p_ij = p_i[j]
            for k in range(3):
                # 批量计算m = a * lambda
                m = a_ij[k] * this_lamda
                # 直接比较3个元素求max,比调用np.max更快
                c = m[0]
                if m[1] > c:
                    c = m[1]
                if m[2] > c:
                    c = m[2]
                # 分步计算logsumexp,避免重复调用exp
                exp0 = math.exp(m[0] - c)
                exp1 = math.exp(m[1] - c)
                exp2 = math.exp(m[2] - c)
                sum_exp = exp0 + exp1 + exp2
                logsumexp = math.log(sum_exp) + c
                # 计算最终概率
                p_ij[k, 0] = math.exp(m[0] - logsumexp)
                p_ij[k, 1] = math.exp(m[1] - logsumexp)
                p_ij[k, 2] = math.exp(m[2] - logsumexp)
    
    return p.transpose(0, 1, 3, 2)

并行化优化版本(多核CPU适用)

@numba.njit(fastmath=True, parallel=True)
def get_p_4d_nb_parallel(a, lamda):
    num_code, num_draw, _, _ = a.shape
    a_trans = a.transpose(0, 1, 3, 2)
    p = np.empty((num_code, num_draw, 3, 3), dtype=a.dtype)
    
    # 用prange并行化外层循环,利用多核CPU
    for i in numba.prange(num_code):
        lamda_i = lamda[i]
        a_i = a_trans[i]
        p_i = p[i]
        for j in range(num_draw):
            this_lamda = lamda_i[j]
            a_ij = a_i[j]
            p_ij = p_i[j]
            for k in range(3):
                m = a_ij[k] * this_lamda
                c = m[0]
                if m[1] > c:
                    c = m[1]
                if m[2] > c:
                    c = m[2]
                exp0 = math.exp(m[0] - c)
                exp1 = math.exp(m[1] - c)
                exp2 = math.exp(m[2] - c)
                sum_exp = exp0 + exp1 + exp2
                logsumexp = math.log(sum_exp) + c
                p_ij[k, 0] = math.exp(m[0] - logsumexp)
                p_ij[k, 1] = math.exp(m[1] - logsumexp)
                p_ij[k, 2] = math.exp(m[2] - logsumexp)
    
    return p.transpose(0, 1, 3, 2)

效果验证

# 测试数据
a = np.ones((112, 1000, 3, 3))
lamda = np.random.uniform(0., 1., size=(112, 1000))

# 验证结果一致性
res_original = get_p_4d(a, lamda)
res_opt = get_p_4d_nb_opt(a, lamda)
res_parallel = get_p_4d_nb_parallel(a, lamda)
print(np.allclose(res_original, res_opt))       # 输出True
print(np.allclose(res_original, res_parallel)) # 输出True

优化效果说明

  • 替换为math模块函数后,fastmath=True会生效,Numba会应用快速数学规则(如忽略NaN检查、近似浮点运算)提升速度
  • 缓存子数组减少了索引计算开销,提升了内存缓存命中率
  • 并行化版本在多核CPU上,能基于原有Numba加速基础再获得2-4倍的性能提升(取决于CPU核心数)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 15:02:44