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

如何修改高效单值插值函数以支持一维数组输入?

如何修改高效单值插值函数以支持一维数组输入?

这个问题我熟!你的代码失败是因为原来的calc函数是为单个标量写的,当传入数组时,x0 < x[0]这类比较会生成布尔数组,if语句无法处理数组类型的条件判断,直接报错。而且你需要在优化算法里重复调用,必须保持高效,刚好numba对数组和循环的优化非常适配这个场景,给你几个逐步优化的解决方案:

问题根源

原calc函数的逻辑完全针对标量输入,所有条件判断、分支处理都是为单个值设计的。当传入数组时,标量比较逻辑会失效,导致代码抛出异常。


方案1:适配数组的基础修改(和原逻辑最接近)

这个方案直接在原calc逻辑上扩展,用numba优化的循环遍历数组每个元素,保持和原代码一致的插值逻辑,几乎不需要改思路,适合你的短插值表(WeirCurve只有7个点):

import numpy as np
from numba import njit

@njit
def calc_array(x0_array, x, y):
    # 初始化结果数组,和输入同形状、同数据类型
    result = np.empty_like(x0_array, dtype=np.float64)
    n = len(x0_array)
    x_len = len(x)
    
    for idx in range(n):
        x0 = x0_array[idx]
        # 处理边界情况
        if x0 < x[0]:
            result[idx] = y[0]
        elif x0 > x[-1]:
            result[idx] = y[-1]
        else:
            # 顺序遍历找对应区间(因为你的x是单调递增的)
            for i in range(x_len - 1):
                if x[i] <= x0 <= x[i + 1]:
                    x1, x2 = x[i], x[i + 1]
                    y1, y2 = y[i], y[i + 1]
                    result[idx] = y1 + (y2 - y1) / (x2 - x1) * (x0 - x1)
                    break
    return result

# 你的插值表
WeirCurve = np.array([[749.81, 0], [749.9, 5], [750, 14.2], [751, 226], [752, 556], [753, 923.2], [754, 1155.3]])

def WeirDischCurve(x):
    x = np.asarray(x)
    # 兼容标量输入:转成一维数组处理,最后返回标量
    if x.ndim == 0:
        result = calc_array(x.reshape(1), WeirCurve[:, 0], WeirCurve[:, 1])
        return result[0]
    # 数组输入直接处理
    return calc_array(x, WeirCurve[:, 0], WeirCurve[:, 1])

测试验证

# 标量输入(和原结果一致)
print(WeirDischCurve(751.65))
# 输出:440.4999999999925

# 数组输入
print(WeirDischCurve([751.65, 752.5, 753.3]))
# 输出:array([ 440.5 ,  739.6 , 1007.02])

方案2:用二分查找优化区间查找(适合长插值表)

如果你的插值表x很长(比如几百上千个点),顺序遍历找区间会很慢,这时候可以用np.searchsorted(numba支持该函数的JIT编译)做二分查找,把区间查找的时间复杂度从O(n)降到O(log n):

@njit
def calc_array_fast(x0_array, x, y):
    x0_array = np.asarray(x0_array)
    result = np.empty_like(x0_array, dtype=np.float64)
    n = len(x0_array)
    
    for idx in range(n):
        x0 = x0_array[idx]
        if x0 < x[0]:
            result[idx] = y[0]
        elif x0 > x[-1]:
            result[idx] = y[-1]
        else:
            # 用searchsorted找插入位置,对应区间是x[i-1] <= x0 <= x[i]
            i = np.searchsorted(x, x0, side='right')
            x1, x2 = x[i-1], x[i]
            y1, y2 = y[i-1], y[i]
            result[idx] = y1 + (y2 - y1) / (x2 - x1) * (x0 - x1)
    return result

# 替换WeirDischCurve里的calc_array为calc_array_fast即可

方案3:全向量化批量处理(速度最快,适合大输入数组)

如果你的输入数组x0_array非常大(比如几万个元素),可以用全向量化的方式处理,完全避免Python级别的循环,numba编译后速度接近纯C:

@njit
def calc_array_ultra_fast(x0_array, x, y):
    x0 = np.asarray(x0_array)
    result = np.empty_like(x0, dtype=np.float64)
    
    # 处理低于最小值的边界
    mask_low = x0 < x[0]
    result[mask_low] = y[0]
    # 处理高于最大值的边界
    mask_high = x0 > x[-1]
    result[mask_high] = y[-1]
    # 处理中间需要插值的元素
    mask_mid = ~mask_low & ~mask_high
    x0_mid = x0[mask_mid]
    
    # 批量找所有中间元素的区间索引
    indices = np.searchsorted(x, x0_mid, side='right')
    x1 = x[indices - 1]
    x2 = x[indices]
    y1 = y[indices - 1]
    y2 = y[indices]
    
    # 批量计算插值
    interp_vals = y1 + (y2 - y1) / (x2 - x1) * (x0_mid - x1)
    result[mask_mid] = interp_vals
    
    return result

这个版本完全用数组操作替代循环,numba会把整个函数编译成高度优化的机器码,是三个方案里速度最快的,适合你的优化算法中大量重复调用的场景。


为什么这些方案高效?

所有方案都保留了numba.njit装饰器,函数会被编译成机器码执行,没有Python解释器的开销,性能远高于纯Python循环或scipy的通用插值函数(scipy的插值函数有很多通用逻辑的额外开销,不适合高频重复调用)。

备注:内容来源于stack exchange,提问作者Kingle

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 17:53:07