如何借助Numba进一步加速多维logsumexp与softmax计算
优化Numba实现的Softmax/LogSumExp计算
现有代码的问题分析
- 过细的循环嵌套:最内层循环逐个处理3个元素,没有利用Numba对数组批量操作的优化能力,反而增加了循环调度开销。
- Numpy函数冗余调用:循环内使用
np.max、np.log等Numpy函数,这类函数在Numba循环内会产生额外调用开销,且阻碍fastmath优化生效——Numba的fastmath对纯Python内置函数或原生操作的优化更显著。 - 内存访问效率低:循环内频繁做单个元素赋值,加上数组转置的额外开销,导致缓存命中率下降。
优化方案与代码实现
核心优化点
- 批量处理固定长度的内层维度(3个元素),避免逐个元素赋值
- 替换Numpy函数为
math模块的原生函数,让fastmath优化生效 - 缓存子数组减少索引计算开销,提升缓存命中率
- 可选:利用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
相关产品推荐
相关产品推荐

