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

如何用Numba Vectorize实现支持单/多维度点的双线性插值函数?

双线性插值函数的Numba vectorize实现问题

需要实现一个双线性插值函数,要求:

  • 输入点数组支持两种形状:(2, n)的二维数组(返回n个插值结果),以及(2,)的一维数组(返回单个插值结果)
  • 尝试用Numba的@vectorize装饰器优雅实现该函数,但未成功

原@njit实现代码

from numba import njit, vectorize, float64
import numpy as np
import math
from numba import prange

@njit
def bilinear_interpolation(points, matrix, axis_0_start, axis_0_step, axis_1_start, axis_1_step):
    res = np.empty(points.shape[1])
    for i in prange(points.shape[1]):
        point0_loc = (points[0, i] - axis_0_start) / axis_0_step
        point1_loc = (points[1, i] - axis_1_start) / axis_1_step

        idx_0l = math.floor(point0_loc)
        idx_0h = idx_0l + 1
        idx_1l = math.floor(point1_loc)
        idx_1h = idx_1l + 1

        mat_hl = matrix[idx_0h, idx_1l]
        mat_ll = matrix[idx_0l, idx_1l]
        mat_hh = matrix[idx_0h, idx_1h]
        mat_lh = matrix[idx_0l, idx_1h]
        
        res[i] =   (mat_ll * (idx_0h - point0_loc) * (idx_1h - point1_loc) +
                    mat_hl * (point0_loc - idx_0l) * (idx_1h - point1_loc) +
                    mat_lh * (idx_0h - point0_loc) * (point1_loc - idx_1l) +
                    mat_hh * (point0_loc - idx_0l) * (point1_loc - idx_1l))
    return res

尝试的@vectorize版本代码

from numba import njit, vectorize, float64
import math

@vectorize([float64(float64, float64, float64[:, :], float64, float64, float64, float64)])
def bilinear_interpolation(point0, point1, matrix, axis_0_start, axis_0_step, axis_1_start, axis_1_step):
    point0_loc = (point0 - axis_0_start) / axis_0_step
    point1_loc = (point1 - axis_1_start) / axis_1_step

    idx_0l = math.floor(point0_loc)
    idx_0h = idx_0l + 1
    idx_1l = math.floor(point1_loc)
    idx_1h = idx_1l + 1

    mat_hl = matrix[idx_0h, idx_1l]
    mat_ll = matrix[idx_0l, idx_1l]
    mat_hh = matrix[idx_0h, idx_1h]
    mat_lh = matrix[idx_0l, idx_1h]
    
    res =   (mat_ll * (idx_0h - point0_loc) * (idx_1h - point1_loc) +
                mat_hl * (point0_loc - idx_0l) * (idx_1h - point1_loc) +
                mat_lh * (idx_0h - point0_loc) * (point1_loc - idx_1l) +
                mat_hh * (point0_loc - idx_0l) * (point1_loc - idx_1l))
    return res

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 12:56:05