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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 05:57:05