Numba中使用np.dot()实现批量矩阵乘法的连续性警告排查
环境版本
- Numba 0.55.1
- Numpy 1.21.5
问题说明
尝试用Numba加速批量矩阵乘法任务,但运行时持续收到np.dot() is faster on contiguous arrays的性能警告。打印所有相关数组的连续性状态,结果均显示为连续,但警告依然存在。
问题代码
import numpy as np import numba as nb def numbaFastMatMult(mat,vec): result = np.zeros_like(vec) for n in nb.prange(vec.shape[0]): result[n,:] = np.dot(vec[n,:], mat[n,:,:]) return result D,N = 10,1000 mat = np.random.normal(0,1,(N,D,D)) vec = np.random.normal(0,1,(N,D)) result = numbaFastMatMult(mat,vec) print(mat.data.contiguous) print(vec.data.contiguous) print(mat[n,:,:].data.contiguous) print(vec[n,:].data.contiguous)
运行输出
连续性打印结果
True True True True
警告信息
NumbaPerformanceWarning: np.dot() is faster on contiguous arrays, called on (array(float64, 1d, C), array(float64, 2d, A))
result[n,:] = np.dot(vec[n,:], mat[n,:,:])
补充说明
- 上述为复现问题的示例代码,实际场景数据量更大,需提升运算速度。
- 已知可通过
np.tensordot解决需求,但希望理解警告出现的根本原因,以便后续参考,现有相关讨论未直接解释此问题。
已尝试的无效方法
- 添加类型装饰器
nb.float64[:,::1](nb.float64[:,:,::1],nb.float64[:,::1]) - 调整批量索引顺序
- 在函数内打印数组连续性状态
问题原因与解决方案
原因解析
虽然整个mat数组是C连续的,但在Numba的JIT编译阶段,切片mat[n,:,:]被识别为**任意连续(A-contiguous)**类型,而非严格的C连续类型。
Numba的类型系统中,A-contiguous仅表示数组内存存储连续,但不保证是标准的C(行优先)或Fortran(列优先)顺序。而np.dot对二维数组的优化依赖于明确的C/Fortran连续布局,因此当传入A-contiguous的二维数组时,会触发性能警告——即使该数组实际在内存中是连续的。
解决方案
显式转换为C连续数组
在切片后用np.ascontiguousarray强制转为C连续,代价是少量内存拷贝,适合小维度场景:result[n,:] = np.dot(vec[n,:], np.ascontiguousarray(mat[n,:,:]))调整数组维度顺序
将批量维度移至最后,如mat = np.random.normal(0,1,(D,D,N)),此时切片mat[:,:,n]会被Numba识别为C连续数组,同时调整vec维度为(D,N)并修改循环逻辑:def numbaFastMatMult(mat,vec): result = np.zeros_like(vec) for n in nb.prange(vec.shape[1]): result[:,n] = np.dot(vec[:,n], mat[:,:,n]) return result使用Numba原生矩阵乘法
替换np.dot为Numba的nb.dot,或手动实现矩阵乘法循环,避免类型推断带来的警告,同时能更好地利用Numba的并行优化:@nb.njit(parallel=True) def numbaFastMatMult(mat,vec): N, D = vec.shape result = np.zeros((N, D), dtype=vec.dtype) for n in nb.prange(N): for i in range(D): tmp = 0.0 for j in range(D): tmp += vec[n,j] * mat[n,j,i] result[n,i] = tmp return result升级Numba版本
新版本Numba对数组连续性的推断逻辑已优化,此类误报警告大概率会被修复。
内容的提问来源于stack exchange,提问作者Andrew Landau

