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

如何更快地沿3D数组轴方向应用一维函数?

问题:加速3D数组沿指定轴的复杂函数计算

我有一个无法直接用np.vectorize或Numba加速的复杂函数,但函数内部包含多个numpy原生向量化操作。需要对大型3D数组沿axis=2(或重塑后沿axis=0)应用这个接收一维输入、输出一维结果的函数。

目前试过并行化np.apply_along_axis、拆分数据和线程的方式,但速度仍不理想;用Dask库在8核机器上的运行时间和当前方案差不多。自己写的并行化代码处理大数组还是慢,尝试用Numba加速但因为调用scipy/numpy原生函数报错无法运行,寻求更高效的实现方法。


现有并行化实现代码

import os, multiprocessing
import scipy
from scipy import stats
import numpy as np

def unpacking_apply_along_axis(all_args):
    (func1d, axis, arr) = all_args
    return np.apply_along_axis(func1d, axis, arr)

def parallel_apply_along_axis_spi(axis, arr):
    chunks = [(precip_2_spi_gh_func, axis, sub_arr)
              for sub_arr in np.array_split(arr, 500)]
    spipool = multiprocessing.Pool(processes=16)
    individual_results = spipool.map(unpacking_apply_along_axis, chunks)
    spipool.close()
    spipool.join()
    
    return np.concatenate(individual_results)

def precip_2_spi_gh_func(ts):
##函数内部有numpy原生向量化操作,但无法整体原生实现##
    ts = np.array(ts)
    ##计算逻辑##
    return zs

调用代码

precip = np.random.uniform(0,300, size=(500,500,43))
# 沿axis=2应用函数
spi = parallel_apply_along_axis_spi(2, precip)

完整函数实现

def precip_2_spi_gh_func(ts):
    ts = np.array(ts)
    normthresh = 160.0
    min_posobs = 12 
    
    # 将降水数据转为向量
    zdim = len(ts)
    pvals  = 0.0  # 正值数量
    logsum = 0.0
    
    pos_ids = np.where(ts > 0)
    pvals = float(len(pos_ids[0]))

    if pvals < min_posobs:
          return np.zeros(zdim) # 非零值数量不足
    
    posave = np.mean(ts[pos_ids])
    logsum = np.sum(np.log(ts[pos_ids]))

    norain_prob = (zdim - pvals) / zdim # 无降水事件占比
    bigA = np.log(posave) - (logsum/pvals)
    
    shape = 0
    scale = 0
    if bigA > 0:
        shape = (1.0+np.sqrt((4.0*bigA/3.0)+1.0)) / (4.0*bigA)
        scale = posave / shape
    
    # shape超过阈值时,用z-score计算SPI
    if shape > normthresh:
        zs = scipy.stats.zscore(ts)

    # shape小于等于阈值时,用伽马分布计算SPI
    if shape <= normthresh: 
        if pvals > 1:
            shape = np.double(shape)
            scale = np.double(scale)
            zs    = np.empty(zdim)
            prob = np.empty(zdim)
            
            for t in range (0,zdim):
                xi  = np.double(ts[t])
                if xi > 0:
                    pxi = scipy.special.gammainc(shape,xi/scale)
                elif xi == 0:
                    pxi = 0
                else:
                    pxi = np.nan

                prob[t] = norain_prob + ((1.0 - norain_prob) * pxi) # 事件概率
                if norain_prob > 0.5:
                    if xi <= 7.0:
                        prob[t] = 0.5

                if np.sum(prob >= 1.0) > 0:
                    prob[np.where(prob >= 1.0)] = 0.99999999                 
                zs[t] = stats.norm.isf(1.0 - prob[t])  

尝试的Numba加速代码

@jit(parallel=True)
def precip_2_spi_gh_func(ts):
    normthresh = 160.0
    min_posobs = 12 
    
    # 将降水数据转为向量
    zdim = len(ts)
    pvals  = 0.0  # 正值数量
    logsum = 0.0
    
    pos_ids = np.where(ts > 0)
    pvals = len(pos_ids[0])

    if pvals < min_posobs:
          return np.zeros(zdim) # 非零值数量不足
    
    posave = np.mean(ts[pos_ids])
    logsum = np.sum(np.log(ts[pos_ids]))

    norain_prob = (zdim - pvals) / zdim # 无降水事件占比
    bigA = np.log(posave) - (logsum/pvals)
    
    shape = 0
    scale = 0
    if bigA > 0:
        shape = (1.0+np.sqrt((4.0*bigA/3.0)+1.0)) / (4.0*bigA)
        scale = posave / shape
    
    # shape超过阈值时,用z-score计算SPI
    if shape > normthresh:
        zs = (ts - np.mean(ts))/np.std(ts)

    # shape小于等于阈值时,用伽马分布计算SPI
    if shape <= normthresh: 
        if pvals > 1:
            zs    = []
            prob = []
            
            for t in range (0,zdim):
                xi  = ts[t]
                if xi > 0:
                    pxi = scipy.special.gammainc(shape,xi/scale)
                elif xi == 0:
                    pxi = 0
                else:
                    pxi = np.nan

                prob[t] = norain_prob + ((1.0 - norain_prob) * pxi) # 事件概率
                if norain_prob > 0.5:
                    if xi <= 7.0:
                        prob[t] = 0.5

                if np.sum(prob >= 1.0) > 0:
                    prob[np.where(prob >= 1.0)] = 0.99999999                 
                zs[t] = stats.norm.isf(1.0 - prob[t])          
    
    zs[np.isnan(zs)] = -9999

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 00:50:55