如何在Cython中为priority_queue传入自定义比较器?
当然可行!你遇到的其实是Cython对接C模板时的常见小细节问题——默认的libcpp.priority_queue包装确实没把自定义比较器的用法直接摆到台面上,但咱们完全能通过自定义C风格的比较器类搞定,还能针对你的场景做高效优化,甚至满足稳定排序的需求。
方法1:自定义C++风格比较器(直接适配priority_queue)
C的priority_queue本身就支持通过第三个模板参数指定自定义比较器,Cython里可以用cdef cppclass定义符合C要求的仿函数,直接传递给模板。
比如你要处理vector[int[:]],按数组的某个元素(比如第一个元素)实现小顶堆(适合合并有序数组的场景),可以这么写:
from libcpp.vector cimport vector from libcpp.queue cimport priority_queue # 定义自定义比较器:按数组第一个元素升序排列(实现小顶堆) cdef cppclass ArrayComparator: bool operator()(const vector[int]& a, const vector[int]& b) noexcept: # 这里可以替换成你需要的任意比较逻辑,比如按指定索引的元素比较 return a[0] > b[0] # 实例化带自定义比较器的优先队列 cdef priority_queue[vector[int], vector[vector[int]], ArrayComparator] pq
用法示例(合并有序数组)
# 准备几个已按目标元素排序的数组 cdef vector[int] arr1 = [1, 3, 5] cdef vector[int] arr2 = [2, 4, 6] cdef vector[int] arr3 = [0, 7, 8] # 将数组推入优先队列 pq.push(arr1) pq.push(arr2) pq.push(arr3) # 合并过程:每次取堆顶最小的数组的首元素,剩余元素重新入堆 cdef vector[int] result while not pq.empty(): cdef vector[int] top_arr = pq.top() pq.pop() result.push_back(top_arr[0]) top_arr.erase(top_arr.begin()) # 移除已取出的首元素 if not top_arr.empty(): pq.push(top_arr)
方法2:用指针/索引优化性能(针对大数组场景)
你提到可以接受转换为缓冲区/指针,那为了避免复制整个数组(尤其是数组很大时),我们可以把数组指针+当前索引存到优先队列里,只操作索引而不复制数据,效率会高很多:
from libcpp.vector cimport vector from libcpp.queue cimport priority_queue # 定义堆元素:保存数组指针、当前读取索引,可选原数组序号(用于稳定排序) cdef cppclass HeapElement: const vector[int]* arr_ptr size_t idx size_t arr_id # 原数组的唯一标识,稳定排序时用 # 构造函数 HeapElement(const vector[int]* ptr, size_t i, size_t aid): arr_ptr = ptr idx = i arr_id = aid # 自定义比较器:先按目标元素升序,元素相等时按原数组序号升序(保证稳定) cdef cppclass HeapElementComparator: bool operator()(const HeapElement& a, const HeapElement& b) noexcept: # 先比较目标元素 if a.arr_ptr->at(a.idx) != b.arr_ptr->at(b.idx): return a.arr_ptr->at(a.idx) > b.arr_ptr->at(b.idx) # 元素相等时,原数组序号小的优先,保证稳定排序 else: return a.arr_id > b.arr_id # 实例化优化后的优先队列 cdef priority_queue[HeapElement, vector[HeapElement], HeapElementComparator] pq
用法示例(高效合并稳定排序)
cdef vector[int] arr1 = [1, 3, 5] cdef vector[int] arr2 = [2, 3, 6] # 包含和arr1相等的元素3 cdef vector[int] arr3 = [0, 7, 8] # 推入堆时传入数组指针、初始索引0、原数组序号 pq.push(HeapElement(&arr1, 0, 0)) pq.push(HeapElement(&arr2, 0, 1)) pq.push(HeapElement(&arr3, 0, 2)) cdef vector[int] result while not pq.empty(): cdef HeapElement top_elem = pq.top() pq.pop() # 取出当前元素 result.push_back(top_elem.arr_ptr->at(top_elem.idx)) # 如果数组还有剩余元素,更新索引后重新入堆 if top_elem.idx + 1 < top_elem.arr_ptr->size(): top_elem.idx += 1 pq.push(top_elem)
这个版本完全避免了数组复制,所有操作都是指针和索引级别的,对于大数组场景性能提升非常明显,同时通过arr_id字段保证了排序的稳定性。
补充:关于“类似argsort”的场景
如果你的需求是对vector[int[:]]做类似argsort的操作(即得到排序后的索引数组),也可以用类似的思路:把数组的索引和目标元素值存入优先队列,最后取出索引即可。比如:
cdef cppclass ArgSortElement: size_t idx int key # 用于排序的目标元素值 cdef cppclass ArgSortComparator: bool operator()(const ArgSortElement& a, const ArgSortElement& b) noexcept: # 按key升序,相等时按原索引升序(稳定) if a.key != b.key: return a.key > b.key else: return a.idx > b.idx # 实例化队列后,把每个数组的索引和目标元素推入,最后取出idx就是argsort结果
内容的提问来源于stack exchange,提问作者Hameer Abbasi
相关产品推荐
相关产品推荐

