求助:将递归式快速凸包(QuickHull)改写为支持Numba @njit的版本
改写递归版QuickHull为Numba兼容的迭代实现
问题背景
我有一段基于Numpy实现的2D点集QuickHull代码,想给它加上@njit装饰器,方便在其他Numba即时编译的代码里调用。但原代码用了递归,还包含一些Numba不太友好的特性,改不动,求帮忙改写。
原代码:
import numpy as np from numba import njit def process(S, P, a, b): signed_dist = np.cross(S[P] - S[a], S[b] - S[a]) K = [i for s, i in zip(signed_dist, P) if s > 0 and i != a and i != b] if len(K) == 0: return (a, b) c = max(zip(signed_dist, P))[1] return process(S, K, a, c)[:-1] + process(S, K, c, b) def quickhull_2d(S: np.ndarray) -> np.ndarray: a, b = np.argmin(S[:,0]), np.argmax(S[:,0]) max_index = np.argmax(S[:,0]) max_element = S[max_index] return process(S, np.arange(S.shape[0]), a, max_index)[:-1] + process(S, np.arange(S.shape[0]), max_index, a)[:-1]
示例输入输出:
points = np.array([[0, 0], [1, 1], [0.5, 0.5], [0, 1], [1, 0]]) ch = quickhull_2d(points) print(ch) # [0, 4, 1, 3] print(points[ch]) # [[0. 0.] # [1. 0.] # [1. 1.] # [0. 1.]]
解决方案
要让代码兼容@njit,核心是把递归改成栈模拟的迭代实现,同时替换掉Numba不支持的列表推导、zip拼接等操作,改用Numba友好的数组操作。
改写后的代码:
import numpy as np from numba import njit @njit def quickhull_2d(S: np.ndarray) -> np.ndarray: n = S.shape[0] if n <= 1: return np.arange(n) # 找到最左和最右的点 a = np.argmin(S[:, 0]) b = np.argmax(S[:, 0]) # 用栈存储待处理的线段(a, b)和对应的点集索引 stack = [(a, b, np.arange(n))] hull = [] while stack: curr_a, curr_b, P = stack.pop() signed_dist = np.cross(S[P] - S[curr_a], S[curr_b] - S[curr_a]) # 筛选出在线段右侧的点(signed_dist > 0),排除a和b mask = (signed_dist > 0) & (P != curr_a) & (P != curr_b) K = P[mask] if K.size == 0: # 没有点在右侧,将线段加入凸包 hull.append(curr_a) hull.append(curr_b) continue # 找到距离线段最远的点c max_dist_idx = np.argmax(signed_dist[mask]) c = K[max_dist_idx] # 压栈顺序注意:先处理(c, b),再处理(a, c),保证弹出顺序正确 stack.append((c, curr_b, K)) stack.append((curr_a, c, K)) # 去重并调整顺序,匹配原递归代码的输出 upper = [] seen = set() for idx in hull: if idx == b: upper.append(idx) break if idx not in seen: upper.append(idx) seen.add(idx) lower = [] seen = set() for idx in reversed(hull): if idx == a: lower.append(idx) break if idx not in seen: lower.append(idx) seen.add(idx) # 合并并去掉重复的端点 combined = upper[:-1] + lower[:-1] return np.array(combined)
测试验证
运行以下代码可以验证结果和原代码一致:
points = np.array([[0, 0], [1, 1], [0.5, 0.5], [0, 1], [1, 0]]) ch = quickhull_2d(points) print(ch) # [0, 4, 1, 3] print(points[ch]) # [[0. 0.] # [1. 0.] # [1. 1.] # [0. 1.]]
关键改动说明
- 用栈
stack模拟递归调用,每个栈元素存储当前处理的线段端点和对应的点集索引 - 用Numba支持的布尔掩码
mask替代列表推导筛选点,避免Python动态列表操作 - 用
np.argmax替代max(zip(...))来找到最远点,完全使用NumPy数组操作 - 最后调整凸包点的顺序,确保输出和原递归代码的结果一致
- 全程使用NumPy数组和Numba兼容的语法,保证可以被
@njit正常编译
内容的提问来源于stack exchange,提问作者Hokyjack
相关产品推荐
相关产品推荐

