如何通过ctypes优化Python调用C++实现的矩阵乘法性能?
优化方案:直接复用Python分配的内存避免数据拷贝
核心优化思路
- 取消C中
new分配内存的操作,改为由Python提前创建输出numpy数组,将其底层内存指针传递给C函数,直接在该内存区域写入计算结果,彻底消除数据拷贝的开销。 - 修正原C++代码中的矩阵索引错误(原索引逻辑不符合行优先存储规则,既导致计算结果错误,也降低了缓存效率)。
- 简化ctypes的类型声明,减少不必要的类型转换步骤。
修改后的代码
1. C++实现(cpp_function.cpp)
编译命令保持不变:g++ -shared -fPIC cpp_function.cpp -o cpp_function.so
#include <iostream> extern "C" { // 直接写入Python预先分配好的内存,无需返回新指针 void mult_matrix(double *a1, double *a2, double *result, size_t a1_h, size_t a1_w, size_t a2_h, size_t a2_w) { // 矩阵乘法:a1(h1,w1) * a2(h2,w2) = result(h1,w2),要求w1=h2 for(size_t i = 0; i < a1_h; i++){ for (size_t j = 0; j < a2_w; j++) { double val = 0.0; // 修正索引:行优先存储下,a1的(i,k)位置是i*a1_w +k for (size_t k = 0; k < a1_w; k++){ val += a1[i * a1_w + k] * a2[k * a2_w + j] ; } // 结果矩阵的(i,j)位置是i*a2_w +j result[i * a2_w + j] = val; } } } }
2. Python调用代码(main.py)
import ctypes import numpy from time import time libmatmult = ctypes.CDLL("./cpp_function.so") # 统一声明numpy数组的指针类型(C连续存储的float64二维数组) ND_POINTER = numpy.ctypeslib.ndpointer(dtype=numpy.float64, ndim=2, flags="C_CONTIGUOUS") # 声明函数参数类型:两个输入矩阵,一个输出矩阵,以及四个尺寸参数 libmatmult.mult_matrix.argtypes = [ ND_POINTER, ND_POINTER, ND_POINTER, ctypes.c_size_t, ctypes.c_size_t, ctypes.c_size_t, ctypes.c_size_t ] # 无返回值 libmatmult.mult_matrix.restype = None def mult_matrix_cpp(a,b): # 提前分配输出数组,和输入数组一样是C连续的float64类型 result_shape = (a.shape[0], b.shape[1]) result = numpy.empty(result_shape, dtype=numpy.float64, order='C') # 直接调用C++函数,传入输入数组、输出数组及各维度尺寸 libmatmult.mult_matrix(a, b, result, a.shape[0], a.shape[1], b.shape[0], b.shape[1]) return result size_a = (300,300) size_b = size_a a = numpy.random.uniform(low=1, high=255, size=size_a).astype(numpy.float64, order='C') b = numpy.random.uniform(low=1, high=255, size=size_b).astype(numpy.float64, order='C') t2 = time() out_cpp = mult_matrix_cpp(a,b) print("cpp time taken:{:.2f} ms".format((time() - t2) * 1000)) t3 = time() out_np = numpy.dot(a,b) print("Numpy dot() time taken:{:.2f} ms".format((time() - t3) * 1000)) # 可选:验证结果正确性 # print(numpy.allclose(out_cpp, out_np))
关键优化点解析
消除数据拷贝开销:
- 原方案中C++分配内存后,Python需要遍历
ret_cpp.contents生成列表再转numpy数组,这两步都是O(n²)的拷贝操作,是性能瓶颈的核心来源。 - 新方案由Python提前分配输出数组,C++直接写入该内存,完全避免了跨语言的内存拷贝。
- 原方案中C++分配内存后,Python需要遍历
修正矩阵索引错误:
- 原代码中
a1[i * a1_h + k]的索引逻辑错误,行优先存储下,每行有a1_w个元素,正确索引应为i * a1_w + k,错误的索引不仅导致计算结果错误,还会破坏缓存局部性,降低计算效率。
- 原代码中
简化ctypes类型声明:
- 复用同一个
ND_POINTER类型声明,减少冗余;明确指定C_CONTIGUOUS确保数组内存布局符合C语言预期,避免隐式转换。
- 复用同一个
内存管理更安全:
- 原方案中C++用
new分配的内存无法被Python的垃圾回收机制自动释放,存在内存泄漏风险;新方案由Python管理所有内存,避免了内存泄漏问题。
- 原方案中C++用
内容的提问来源于stack exchange,提问作者Devil
相关产品推荐
相关产品推荐

