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

求助:将递归式快速凸包(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 05:30:42