Numba njit装饰numpy实现的correlate函数报TypingError求解
问题描述
我编写了如下correlate互相关函数:
@numba.njit def correlate(self, mat, filter, type) -> np.ndarray: def apply_filter(mat, filter, point): point = (max(0, point[0]), max(0, point[1])) end_point = (min(mat.shape[0], point[0] + filter.shape[0]), min(mat.shape[1], point[1] + filter.shape[1])) area = mat[point[0]:end_point[0], point[1]:end_point[1]] if filter.shape != area.shape: # filter = np.resize(filter, area.shape) filter = filter[area.shape[0] - 1, area.shape[1] - 1] result = np.multiply(area, filter) result = np.sum(result) return result if type == "valid": new_mat_size = (mat.shape[0] - filter.shape[0] + 1, mat.shape[1] - filter.shape[0] + 1) elif type == "full": new_mat_size = (mat.shape[0] + filter.shape[0] - 1, mat.shape[1] + filter.shape[1] - 1) f_mat = np.zeros(new_mat_size) for x in range(new_mat_size[0]): for y in range(new_mat_size[1]): if type == "valid": f_mat[x, y] = apply_filter(mat, filter, (x, y)) elif type == "full": f_mat[x, y] = apply_filter(mat, filter, (x - filter.shape[0] + 1, y - filter.shape[1] + 1)) return f_mat
该函数在不添加@numba.njit装饰器时可正常运行,添加装饰器后抛出如下核心错误:
numba.core.errors.TypingError: Failed in nopython mode pipeline (step: nopython frontend) Cannot unify float64 and array(float64, 2d, C) for 'closure__locals__apply_filter_v3_filter_2.2' ... This error may have been caused by the following argument(s): - argument 0: Cannot determine Numba type of <class 'neuralgNeccessities.Layers.convolutional_layer.Convolutional'>
错误原因
报错全部来自Numba nopython模式的静态编译规则限制,附带一处隐藏逻辑bug:
- 自定义类参数无法识别:
@numba.njit被加在了类的实例方法上,方法第一个参数是类实例self,Numba无法对自定义Python类做类型推断,直接触发参数类型解析失败。且该函数内部完全没有用到self的任何属性或方法,根本不需要定义为实例方法。 - 变量类型冲突:内部函数
apply_filter中,filter初始传入为2维浮点数组,但进入形状不匹配分支时,filter = filter[area.shape[0] - 1, area.shape[1] - 1]的写法是取数组中的单个标量值,Numba编译时要求同一个变量不能同时持有数组和标量两种类型,直接触发类型统一失败。同时这行代码本身逻辑错误:注释掉的np.resize是调整数组形状,现有写法是取单个元素,纯Python环境下走到该分支计算结果也不符合互相关逻辑。 - 隐藏尺寸计算bug:计算
valid模式输出尺寸时,宽度维度错误使用了filter.shape[0](卷积核高度),传入非正方形卷积核时会直接输出错误尺寸。
修复方案
- 把互相关逻辑从类实例方法中抽离,写成不依赖
self的独立函数后再加@numba.njit装饰器,调用时直接传入mat、filter、计算模式三个参数即可。注意不要用type作为变量名,这是Python内置关键字,替换为mode更稳妥。 - 修正filter形状不匹配时的处理逻辑:full模式下做边缘互相关时,对filter做切片和截取的area区域形状对齐即可,不要在分支中给传入的
filter参数重新赋值,避免类型冲突。 - 修正valid模式的输出尺寸计算逻辑,宽度维度使用卷积核宽度
filter.shape[1]计算。
修复后的可运行代码参考:
import numba import numpy as np @numba.njit def _apply_filter(mat, kernel, point): point = (max(0, point[0]), max(0, point[1])) end_point = (min(mat.shape[0], point[0] + kernel.shape[0]), min(mat.shape[1], point[1] + kernel.shape[1])) area = mat[point[0]:end_point[0], point[1]:end_point[1]] # 裁剪卷积核到和邻域相同形状,适配边缘位置计算 kernel_cropped = kernel[:area.shape[0], :area.shape[1]] return np.sum(area * kernel_cropped) @numba.njit def correlate(mat, kernel, mode) -> np.ndarray: kh, kw = kernel.shape mh, mw = mat.shape if mode == "valid": new_h = mh - kh + 1 new_w = mw - kw + 1 f_mat = np.zeros((new_h, new_w), dtype=mat.dtype) for x in range(new_h): for y in range(new_w): f_mat[x, y] = _apply_filter(mat, kernel, (x, y)) elif mode == "full": new_h = mh + kh - 1 new_w = mw + kw - 1 f_mat = np.zeros((new_h, new_w), dtype=mat.dtype) for x in range(new_h): for y in range(new_w): f_mat[x, y] = _apply_filter(mat, kernel, (x - kh + 1, y - kw + 1)) return f_mat
额外优化点:提前把卷积核、输入矩阵的高宽存为局部变量,减少Numba编译时的属性查找开销;输出矩阵指定和输入一致的dtype,避免默认float64带来的不必要类型转换;把内部计算逻辑抽为平级的njit函数,减少闭包带来的编译开销。
内容的提问来源于stack exchange,提问作者eirikg
相关产品推荐
相关产品推荐

