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

在Cython中为Shapely对象使用何种类型声明以加速代码?

问题描述

我是Cython新手,这可能是一个入门问题。我有一段Python代码(仅单个函数),在脚本中会被调用数千次,因此哪怕小幅提速也能大幅缩短脚本运行时间。我尝试用Cython来加速,通过教程了解到第一步是为函数添加正确的C类型声明,但由于该函数使用Shapely对象,我不知道应使用何种类型。

原始Python代码

import numpy as np
from shapely.geometry import LineString, Point
from shapely.linear import line_locate_point
from shapely.strtree import STRtree

def xy2sd(traj, centerline):
    '''
    Transforms a set of points into the frenet frame given by centerline
    traj: np.ndarray of shape [N, 2] where N is the number of points in the trajectory
    centerline: np.ndarray of shape [M, 2] where M is the number of points in the centerline
    '''
    line = LineString(centerline)
    s_orig = line_locate_point(line, Point([0, 0]))

    s = np.zeros(traj.shape[0])
    d = np.zeros(traj.shape[0])

    nth_entry = 100
    tree = STRtree([Point(x) for id, x in enumerate(line.coords) if id % nth_entry == 0])

    for i in range(traj.shape[0]):
        pnt = Point(traj[i])
        s[i] = line.line_locate_point(pnt)
        d[i] = pnt.distance(line)

        min_idx = tree.nearest(pnt)*nth_entry

        pnt_min_dist = centerline[min_idx]
        pnt_behind = centerline[min_idx-1]

        sign = (pnt_behind[0]-pnt_min_dist[0])*(pnt.y-pnt_min_dist[1]) - \
            (pnt_behind[1]-pnt_min_dist[1])*(pnt.x-pnt_min_dist[0])
        sign = 1 if sign >= 0 else -1 if sign < 0 else 0

        d[i] = sign*d[i]

    sd_coords = np.zeros(traj.shape, dtype=np.float32)
    sd_coords[:, 0] = s - s_orig
    sd_coords[:, 1] = d

    return sd_coords

修改后的Cython代码

cimport cython
cimport numpy as np

import numpy as np
from shapely.geometry import LineString, Point
from shapely.linear import line_locate_point
from shapely.strtree import STRtree

def xy2sd(traj, centerline):
    '''
    Transforms a set of points into the frenet frame given by centerline
    traj: np.ndarray of shape [N, 2] where N is the number of points in the trajectory
    centerline: np.ndarray of shape [M, 2] where M is the number of points in the centerline
    '''
    line = LineString(centerline)
    s_orig = line_locate_point(line, Point([0, 0]))

    cdef np.ndarray s = np.zeros(traj.shape[0])
    cdef np.ndarray d = np.zeros(traj.shape[0])

    cdef int nth_entry = 100
    tree = STRtree([Point(x) for id, x in enumerate(line.coords) if id % nth_entry == 0])

    cdef int i
    cdef int min_idx
    cdef np.ndarray pnt_min_dist
    cdef np.ndarray pnt_behind
    cdef double cross_prod
    cdef int sign
    for i in range(traj.shape[0]):
        pnt = Point(traj[i])
        s[i] = line.line_locate_point(pnt)
        d[i] = pnt.distance(line)

        min_idx = tree.nearest(pnt)*nth_entry

        pnt_min_dist = centerline[min_idx]
        pnt_behind = centerline[min_idx-1]

        cross_prod = (pnt_behind[0]-pnt_min_dist[0])*(pnt.y-pnt_min_dist[1]) - \
            (pnt_behind[1]-pnt_min_dist[1])*(pnt.x-pnt_min_dist[0])
        sign = 1 if cross_prod >= 0 else -1 if cross_prod < 0 else 0

        d[i] = sign*d[i]

    cdef np.ndarray sd_coords = np.zeros(traj.shape, dtype=np.float32)
    sd_coords[:, 0] = s - s_orig
    sd_coords[:, 1] = d

    return sd_coords

我能成功编译这段代码,但未发现任何提速效果。我猜测这是因为缺少对Shapely对象line、tree和pnt的类型声明。请问我应为这些对象使用何种类型?仅添加类型声明能否带来显著提速,或是还有其他可优化的地方?


解决方案

1. 关于Shapely对象的Cython类型声明

Shapely的Python对象基于GEOS库的C扩展实现,但没有公开可供Cython直接使用的静态类型定义,无法给LineString、Point、STRtree这类对象做静态类型声明。即使声明为object类型,也不会带来任何提速——这些对象的方法调用本质还是Python层面的交互,Cython无法优化这类外部库的Python API调用。

2. 当前Cython代码无提速的原因

你的循环中90%以上的耗时都集中在Shapely的Python方法调用上:line.line_locate_point()、pnt.distance()、tree.nearest()。你添加的numpy数组和基础类型声明,仅优化了极少量循环变量和数组操作,对整体耗时几乎没有影响。

3. 有效的优化方向

(1)用Shapely 2.0+的向量化API替代循环

Shapely 2.0合并了pygeos的功能,支持向量化操作,可一次性处理整个traj数组,避免循环内的Python对象创建和方法调用:

from shapely.geometry import MultiPoint

def xy2sd_vectorized(traj, centerline):
    line = LineString(centerline)
    s_orig = line_locate_point(line, Point([0, 0]))
    
    # 向量化创建多点对象
    traj_points = MultiPoint(traj)
    # 批量计算s和d
    s = np.array([line.line_locate_point(p) for p in traj_points])
    d = np.array(traj_points.distance(line))
    
    # 批量处理最近点查找
    nth_entry = 100
    sample_points = [Point(x) for id, x in enumerate(line.coords) if id % nth_entry == 0]
    tree = STRtree(sample_points)
    nearest_indices = np.array([tree.nearest(p) for p in traj_points]) * nth_entry
    
    # 向量化计算符号
    pnt_min_dist = centerline[nearest_indices]
    pnt_behind = centerline[nearest_indices - 1]
    cross_prod = (pnt_behind[:,0] - pnt_min_dist[:,0]) * (traj[:,1] - pnt_min_dist[:,1]) - \
                 (pnt_behind[:,1] - pnt_min_dist[:,1]) * (traj[:,0] - pnt_min_dist[:,0])
    sign = np.where(cross_prod >= 0, 1, np.where(cross_prod < 0, -1, 0))
    d *= sign
    
    sd_coords = np.zeros(traj.shape, dtype=np.float32)
    sd_coords[:,0] = s - s_orig
    sd_coords[:,1] = d
    return sd_coords

向量化操作能大幅减少Python循环的开销,比Cython优化这类Python调用密集的代码效果更明显。

(2)减少Shapely对象的创建

循环中每次创建Point(traj[i])是较大开销,尽量直接从numpy数组读取坐标——比如符号计算部分,用traj[i,0]、traj[i,1]代替pnt.x、pnt.y,可节省对象属性访问的时间。

(3)预计算可复用资源

如果centerline在多次函数调用中固定,可以提前创建LineString对象、STRtree采样树并传入函数,避免每次调用重复初始化,能节省大量时间。

(4)直接调用GEOS的C API(进阶)

如果追求极致性能,可在Cython中直接调用GEOS的C函数,跳过Shapely的Python层封装。这需要熟悉GEOS的C API并手动处理内存管理,门槛较高,但能获得最大提速。


内容的提问来源于stack exchange,提问作者M_M246

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 16:44:54