我的Cython纯Python模式代码性能弱于Numpy,求优化指导
问题:Cython实现欧氏距离计算性能不如Numpy的优化疑问
我正在学习Cython,为后续实现C与Python的集成做准备。本次练习的功能是:
- 从文件读取包含时间步的两组3D坐标长列表
- 计算每个时间步两点间的欧氏距离,输出为numpy数组
我用Cython的Pure Python模式写了代码,具体实现如下:
# computer.py import cython import numpy as np if cython.compiled: print("Yep, I'm compiled.") from cython.cimports.libc.math import sqrt else: print("Just a lowly interpreted script.") from math import sqrt @cython.boundscheck(False) @cython.wraparound(False) @cython.cfunc def compute_distance_cy(x1: cython.float, y1: cython.float, z1: cython.float, x2: cython.float, y2: cython.float, z2: cython.float): return sqrt(sum(((x1 - x2) ** 2.0, (y1 - y2) ** 2.0, (z1 - z2) ** 2.0))) @cython.boundscheck(False) @cython.wraparound(False) @cython.ccall def compute_distances_pure(points: cython.double[:, :]): # get the maximum dimensions of the array x_max: cython.size_t = points.shape[0] y_max: cython.size_t = points.shape[1] # create memoryviews of the single points view2d: cython.double[:, :] = points view1d: cython.double[:] # create memoryviews of the results result = np.zeros(x_max, dtype=np.double) result2dview: cython.double[:] = result # access the memoryview by way of our constrained indexes x: cython.size_t for x in range(x_max): view1d = view2d[x, :] result2dview[x] = compute_distance_cy( view1d[0], view1d[1], view1d[2], view1d[3], view1d[4], view1d[5]) return result
调用代码:
... def pure_python_mode(points): return computer.compute_distances_pure(points) def do_it_in_numpy(points): return np.sqrt((points[:, 0] - points[:, 3])**2 + (points[:, 1] - points[:, 4])**2 + (points[:, 2] - points[:, 5])**2) points_ndarray = np.array(points_list, dtype=np.double) points_distance_array_from_cython = pure_python_mode(points_ndarray) points_distance_array_from_numpy = do_it_in_numpy(points_ndarray)
计时结果:
Function pure_python_mode Took 0.0439 seconds Function do_it_in_numpy Took 0.0183 seconds
目前我的Cython代码性能比Numpy慢4-8倍,虽然满足当前需求,但想确认是否存在代码问题,或是该场景下性能已达上限。
优化建议
- 避免Python层面的
sum()调用compute_distance_cy里的sum()会引入Python函数调用开销,直接展开计算并匹配参数类型:
@cython.boundscheck(False) @cython.wraparound(False) @cython.cfunc def compute_distance_cy(x1: cython.double, y1: cython.double, z1: cython.double, x2: cython.double, y2: cython.double, z2: cython.double): dx = x1 - x2 dy = y1 - y2 dz = z1 - z2 return sqrt(dx*dx + dy*dy + dz*dz)
将参数类型改为cython.double,和输入数组的double类型一致,避免不必要的类型转换。
- 省去一维视图创建开销
每次循环创建view1d = view2d[x, :]会产生额外开销,直接通过二维内存视图访问元素:
@cython.boundscheck(False) @cython.wraparound(False) @cython.ccall def compute_distances_pure(points: cython.double[:, :]): x_max: cython.size_t = points.shape[0] result = np.zeros(x_max, dtype=np.double) result_view: cython.double[:] = result x: cython.size_t for x in range(x_max): dx = points[x, 0] - points[x, 3] dy = points[x, 1] - points[x, 4] dz = points[x, 2] - points[x, 5] result_view[x] = sqrt(dx*dx + dy*dy + dz*dz) return result
直接在循环内完成计算,省去函数调用和视图创建的额外消耗。
- 开启编译级优化
在编译脚本中添加O3优化选项,让编译器做代码优化(如循环展开、指令并行):
from setuptools import setup from Cython.Build import cythonize import numpy as np setup( ext_modules=cythonize("computer.py", compiler_directives={'language_level': 3}, annotate=True), include_dirs=[np.get_include()], extra_compile_args=["-O3"], )
- 大数据量下尝试并行计算
如果数据规模较大,可通过OpenMP启用多线程并行,编译时添加对应参数:
from setuptools import setup from Cython.Build import cythonize import numpy as np setup( ext_modules=cythonize("computer.py", compiler_directives={'language_level': 3, 'boundscheck': False, 'wraparound': False}), include_dirs=[np.get_include()], extra_compile_args=["-O3", "-fopenmp"], extra_link_args=["-fopenmp"], )
修改循环为并行模式:
from cython.cimports.openmp import omp @cython.boundscheck(False) @cython.wraparound(False) @cython.ccall def compute_distances_pure(points: cython.double[:, :]): x_max: cython.size_t = points.shape[0] result = np.zeros(x_max, dtype=np.double) result_view: cython.double[:] = result x: cython.size_t for x in cython.parallel.prange(x_max, nogil=True): dx = points[x, 0] - points[x, 3] dy = points[x, 1] - points[x, 4] dz = points[x, 2] - points[x, 5] result_view[x] = sqrt(dx*dx + dy*dy + dz*dz) return result
注意小数据量下并行开销可能超过收益,需根据实际规模测试。
内容的提问来源于stack exchange,提问作者largehadroncollider
相关产品推荐
相关产品推荐

