Python海量计算嵌套for循环加速方案及C/C++嵌入咨询
针对海量计算Python代码的优化方案及C/C++嵌入建议
首先必须明确:原代码的嵌套循环在n=10万时会产生1e10次迭代,Python原生循环根本无法在合理时间内完成,必须通过以下方式优化。
一、先修正原代码的明显问题
原代码里的判断条件if d >= tresh and d <= tresh完全等价于d == tresh,但浮点数直接判等极易出现精度误差,建议改为判断距离的平方等于阈值的平方(避免开根号,还能规避精度问题),或者用近似相等判断abs(d - tresh) < 1e-9。
二、Python内的优化方案(无需C/C++经验)
1. NumPy向量化重构(最推荐,速度提升几个数量级)
NumPy的底层是C实现的,向量化运算能彻底消除Python嵌套循环的开销,直接处理数组级别的运算。重构后的代码如下:
import numpy as np tresh = 30 tresh_sq = tresh ** 2 # 预计算阈值平方,避免重复计算 n = 100000 # 直接生成numpy数组,替代循环append,速度快数倍 x1 = np.random.randint(0, 101, size=n) x2 = np.random.randint(0, 101, size=n) y1 = np.random.randint(0, 101, size=n) y2 = np.random.randint(0, 101, size=n) def calc(): # 用广播机制计算所有点对的坐标差平方和 dx = y1[:, np.newaxis] - x1 dy = y2[:, np.newaxis] - x2 dist_sq = dx ** 2 + dy ** 2 # 筛选出符合条件的点对索引(用isclose避免浮点数精度问题) mask = np.isclose(dist_sq, tresh_sq, atol=1e-9) m_indices, n_indices = np.where(mask) # 批量计算结果 a = (y1[m_indices] + x1[n_indices]) / 2.0 b = (y2[m_indices] + x2[n_indices]) / 2.0 c = np.sqrt(dist_sq[mask]) # 或直接用tresh,因为已满足条件 return a, b, c a, b, c = calc()
2. Numba JIT编译(改动极小,速度接近C)
如果你不想重构代码,用Numba的即时编译可以直接把Python循环编译成机器码,只需要给函数加一个装饰器:
import math import random from numba import jit tresh = 30 tresh_sq = tresh ** 2 n = 100000 # 用列表推导生成初始数据,比循环append快 x1 = [random.randint(0, 100) for _ in range(n)] x2 = [random.randint(0, 100) for _ in range(n)] y1 = [random.randint(0, 100) for _ in range(n)] y2 = [random.randint(0, 100) for _ in range(n)] @jit(nopython=True) # 开启nopython模式,编译为纯机器码,速度最快 def calc(x1, x2, y1, y2, tresh_sq): a = [] b = [] c = [] x1_len = len(x1) y1_len = len(y1) for n_idx in range(x1_len): for m_idx in range(y1_len): dx = y1[m_idx] - x1[n_idx] dy = y2[m_idx] - x2[n_idx] dist_sq = dx ** 2 + dy ** 2 if dist_sq == tresh_sq: a.append((y1[m_idx] + x1[n_idx]) / 2.0) b.append((y2[m_idx] + x2[n_idx]) / 2.0) c.append(math.sqrt(dist_sq)) return a, b, c a, b, c = calc(x1, x2, y1, y2, tresh_sq)
注:第一次调用函数会有编译开销,后续调用速度极快。
3. 基础小优化(配合上述方案使用)
- 避免全局变量:把
x1、x2等变量作为函数参数传入,减少全局查找开销。 - 预计算重复值:比如
tresh_sq,避免在循环里重复计算平方和开根号。 - 用列表推导生成初始数据:比循环
append快得多。
三、引入C/C++的嵌入建议(当Python优化仍不够时)
如果上述Python优化仍达不到性能要求,可以用以下几种方式嵌入C/C代码,无需深入掌握复杂的C/C语法:
1. Cython(最容易上手,语法接近Python)
Cython是Python的超集,允许给变量加类型声明,编译成C扩展后速度接近纯C:
- 安装Cython:
pip install cython - 编写
calc.pyx文件(给原代码加类型标注):
import math cimport cython # 关闭边界检查和负索引,提升速度 @cython.boundscheck(False) @cython.wraparound(False) def calc(list x1, list x2, list y1, list y2, int tresh_sq): cdef list a = [] cdef list b = [] cdef list c = [] cdef int x1_len = len(x1) cdef int y1_len = len(y1) cdef int n_idx, m_idx cdef int dx, dy cdef int dist_sq for n_idx in range(x1_len): for m_idx in range(y1_len): dx = y1[m_idx] - x1[n_idx] dy = y2[m_idx] - x2[n_idx] dist_sq = dx * dx + dy * dy if dist_sq == tresh_sq: a.append((y1[m_idx] + x1[n_idx]) / 2.0) b.append((y2[m_idx] + x2[n_idx]) / 2.0) c.append(math.sqrt(dist_sq)) return a, b, c
- 编写
setup.py编译成扩展:
from setuptools import setup from Cython.Build import cythonize setup( ext_modules = cythonize("calc.pyx") )
- 编译:
python setup.py build_ext --inplace,之后就可以像普通Python模块一样导入calc。
2. Pybind11(适合简单C++代码包装)
Pybind11可以把C++函数直接包装成Python可调用的函数,步骤简单:
- 安装pybind11:
pip install pybind11 - 编写C++代码
calc.cpp:
#include <pybind11/pybind11.h> #include <vector> #include <cmath> namespace py = pybind11; py::tuple calc(const std::vector<int>& x1, const std::vector<int>& x2, const std::vector<int>& y1, const std::vector<int>& y2, int tresh_sq) { std::vector<double> a, b, c; int x1_len = x1.size(); int y1_len = y1.size(); for (int n_idx = 0; n_idx < x1_len; ++n_idx) { for (int m_idx = 0; m_idx < y1_len; ++m_idx) { int dx = y1[m_idx] - x1[n_idx]; int dy = y2[m_idx] - x2[n_idx]; int dist_sq = dx * dx + dy * dy; if (dist_sq == tresh_sq) { a.push_back((y1[m_idx] + x1[n_idx]) / 2.0); b.push_back((y2[m_idx] + x2[n_idx]) / 2.0); c.push_back(std::sqrt(dist_sq)); } } } return py::make_tuple(a, b, c); } PYBIND11_MODULE(calc, m) { m.def("calc", &calc, "Calculate matching point pairs"); }
- 编译成扩展(通过setup.py或直接用编译器命令),之后导入使用即可。
3. ctypes(调用已编译的C动态库)
如果你已经有编译好的C动态库,可以用ctypes直接调用,无需修改Python代码结构:
- 编写C代码
calc.c:
#include <stdlib.h> #include <math.h> // 定义返回结果的结构体 typedef struct { double* a; double* b; double* c; int count; } Result; Result calc(int* x1, int* x2, int* y1, int* y2, int n, int tresh_sq) { Result res; res.count = 0; // 先统计符合条件的点对数量 for (int i = 0; i < n; i++) { for (int j = 0; j < n; j++) { int dx = y1[j] - x1[i]; int dy = y2[j] - x2[i]; if (dx*dx + dy*dy == tresh_sq) { res.count++; } } } // 分配内存存储结果 res.a = (double*)malloc(res.count * sizeof(double)); res.b = (double*)malloc(res.count * sizeof(double)); res.c = (double*)malloc(res.count * sizeof(double)); int idx = 0; // 再次遍历计算结果 for (int i = 0; i < n; i++) { for (int j = 0; j < n; j++) { int dx = y1[j] - x1[i]; int dy = y2[j] - x2[i]; if (dx*dx + dy*dy == tresh_sq) { res.a[idx] = (y1[j] + x1[i]) / 2.0; res.b[idx] = (y2[j] + x2[i]) / 2.0; res.c[idx] = sqrt(dx*dx + dy*dy); idx++; } } } return res; } // 用于释放内存的函数 void free_result(Result* res) { free(res->a); free(res->b); free(res->c); }
- 编译成动态库:Linux下
gcc -shared -fPIC calc.c -o calc.so,Windows下gcc -shared -fPIC calc.c -o calc.dll - Python中调用:
import ctypes import numpy as np # 加载动态库 lib = ctypes.CDLL('./calc.so') # 定义结构体类型 class Result(ctypes.Structure): _fields_ = [("a", ctypes.POINTER(ctypes.c_double)), ("b", ctypes.POINTER(ctypes.c_double)), ("c", ctypes.POINTER(ctypes.c_double)), ("count", ctypes.c_int)] # 声明函数的参数和返回类型 lib.calc.restype = Result lib.calc.argtypes = [ctypes.POINTER(ctypes.c_int), ctypes.POINTER(ctypes.c_int), ctypes.POINTER(ctypes.c_int), ctypes.POINTER(ctypes.c_int), ctypes.c_int, ctypes.c_int] lib.free_result.argtypes = [ctypes.POINTER(Result)] n = 100000 tresh_sq = 30 ** 2 # 生成numpy数组,方便转换为C指针 x1 = np.random.randint(0, 101, size=n, dtype=np.int32) x2 = np.random.randint(0, 101, size=n, dtype=np.int32) y1 = np.random.randint(0, 101, size=n, dtype=np.int32) y2 = np.random.randint(0, 101, size=n, dtype=np.int32) # 调用C函数 res = lib.calc(x1.ctypes.data_as(ctypes.POINTER(ctypes.c_int)), x2.ctypes.data_as(ctypes.POINTER(ctypes.c_int)), y1.ctypes.data_as(ctypes.POINTER(ctypes.c_int)), y2.ctypes.data_as(ctypes.POINTER(ctypes.c_int)), n, tresh_sq) # 转换为Python列表 a = [res.a[i] for i in range(res.count)] b = [res.b[i] for i in range(res.count)] c = [res.c[i] for i in range(res.count)] # 释放C分配的内存 lib.free_result(ctypes.byref(res))
内容的提问来源于stack exchange,提问作者Per Helge Semb
相关产品推荐
相关产品推荐

