【问题标题】:Cython: How to speed up recursive functions?Cython:如何加速递归函数?
【发布时间】:2020-06-20 13:34:17
【问题描述】:

我正在 cython 中实现一个段树,并将其与 python 实现进行比较。

cython 版本似乎只快 1.5 倍,我想让它更快。

可以假设这两种实现都是正确的。

这是 cython 代码:

# distutils: language = c++
from libcpp.vector cimport vector

cdef struct Result:
    int range_sum  
    int range_min 
    int range_max



cdef class SegmentTree:
    cdef vector[int] nums
    cdef vector[Result] tree 

    def __init__(self, vector[int] nums):
        self.nums = nums
        self.tree.resize(4 * len(nums)) #just a safe upper bound 
        self._build(1, 0, len(nums)-1)

    cdef Result _build(self, int index, int left, int right):
        cdef Result result

        if left == right:
            value = self.nums[left]
            result.range_max, result.range_min, result.range_sum = value, value, value 
            self.tree[index] = result
            return self.tree[index]
        else:
            mid = (left+right)//2
            left_range_result = self._build(index*2, left, mid)
            right_range_result = self._build(index*2+1, mid+1, right)
            self.tree[index] = self.combine_range_results(left_range_result, right_range_result)
            return self.tree[index]

    cdef Result range_query(self, int query_i, int query_j):
        return self._range_query(query_i, query_j, 0, len(self.nums)-1, 1)

    cdef Result _range_query(self, int query_i, int query_j, int current_i, int current_j, int index):
        if current_i == query_i and current_j == query_j:
            return self.tree[index]
        else:
            mid = (current_i + current_j)//2 
            if query_j <= mid:
                return self._range_query(query_i, query_j, current_i, mid, index*2)
            elif mid < query_i:
                return self._range_query(query_i, query_j, mid+1, current_j, index*2+1 )  
            else:
                left_range_result = self._range_query(query_i, mid, current_i, mid, index*2)
                right_range_result = self._range_query(mid+1, query_j, mid+1, current_j, index*2+1)
                return self.combine_range_results(left_range_result, right_range_result)


    cpdef int range_sum(self, int query_i, int query_j):
        return self.range_query(query_i, query_j).range_sum 
    cpdef int range_min(self, int query_i, int query_j):
        return self.range_query(query_i, query_j).range_min
    cpdef int range_max(self, int query_i, int query_j):
        return self.range_query(query_i, query_j).range_max

    cpdef void  update(self, int i, int new_value):
        self._update(i, new_value, 1, 0, len(self.nums)-1)

    cdef Result _update(self, int i, int new_value, int index, int left, int right):
        if left == right == i:
            self.tree[index] = [new_value, new_value, new_value]
            return self.tree[index]
        if left == right:
            return self.tree[index]
        mid = (left+right)//2 
        left_range_result = self._update(i, new_value, index*2, left, mid)
        right_range_result = self._update(i, new_value, index*2+1, mid+1, right)
        self.tree[index] = self.combine_range_results(left_range_result, right_range_result)
        return self.tree[index]

    cdef Result combine_range_results(self, Result r1, Result r2):
        cdef Result result;
        result.range_min = min(r1.range_min, r2.range_min)
        result.range_max = max(r1.range_max, r2.range_max)
        result.range_sum = r1.range_sum + r2.range_sum
        return result 
        

这是python版本:




class PurePythonSegmentTree:
    def __init__(self, nums):
        self.nums = nums
        self.tree = [0] * (len(nums) * 4)
        self._build(1, 0, len(nums) - 1)

    def _build(self, index, left, right):
        if left == right:
            value = self.nums[left]
            self.tree[index] = (value, value, value)
            return self.tree[index]
        else:
            mid = (left + right) // 2
            left_range_result = self._build(index * 2, left, mid)
            right_range_result = self._build(index * 2 + 1, mid + 1, right)
            self.tree[index] = self._combine_range_results(
                left_range_result, right_range_result)
            return self.tree[index]

    def range_query(self, query_i, query_j):
        return self._range_query(query_i, query_j, 0, len(self.nums) - 1, 1)

    def _range_query(self, query_i, query_j, current_i, current_j, index):
        if current_i == query_i and current_j == query_j:
            return self.tree[index]
        else:
            mid = (current_i + current_j) // 2
            if query_j <= mid:
                return self._range_query(query_i, query_j, current_i, mid,
                                         index * 2)
            elif mid < query_i:
                return self._range_query(query_i, query_j, mid + 1, current_j,
                                         index * 2 + 1)
            else:
                left_range_result = self._range_query(query_i, mid, current_i,
                                                      mid, index * 2)
                right_range_result = self._range_query(mid + 1, query_j,
                                                       mid + 1, current_j,
                                                       index * 2 + 1)
                return self._combine_range_results(left_range_result,
                                                   right_range_result)

    def range_sum(self, query_i, query_j):
        return self.range_query(query_i, query_j)[0]

    def range_min(self, query_i, query_j):
        return self.range_query(query_i, query_j)[1]

    def range_max(self, query_i, query_j):
        return self.range_query(query_i, query_j)[2]

    def _combine_range_results(self, r1, r2):
        return (r1[0] + r2[0], min(r1[1], r2[1]), max(r1[2], r2[2]))


基准测试代码:

import pytest
from segment_tree import SegmentTree

def _test_all_ranges(nums, correct_fn, test_fn, threshold=float("inf")):
    count = 0
    for i in range(len(nums)):
        for j in range(i + 1, len(nums)):
            if count > threshold:
                break
            expected = correct_fn(nums[i:j + 1])
            actual = test_fn(i, j)
            assert actual == expected
            count += 1


def test_cython_tree_speed(benchmark):
    nums = [i for i in range(1000)]

    @benchmark
    def foo():
        s = SegmentTree(nums)
        _test_all_ranges(nums, max, s.range_max, 20)


def test_python_tree_speed(benchmark):
    nums = [i for i in range(1000)]

    @benchmark
    def foo():
        s = PurePythonSegmentTree(nums)
        _test_all_ranges(nums, max, s.range_max, 20)

统计数据:

-------------------------------------------------------------------------------------------- benchmark: 2 tests --------------------------------------------------------------------------------------------
Name (time in us)                 Min                   Max                  Mean              StdDev                Median                IQR            Outliers         OPS            Rounds  Iterations
------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------
test_cython_tree_speed       708.0450 (1.0)      1,534.6150 (1.0)        739.7052 (1.0)       59.9436 (1.0)        717.7565 (1.0)      21.0070 (1.0)       116;200  1,351.8900 (1.0)        1290           1
test_python_tree_speed     1,625.1940 (2.30)     2,676.9020 (1.74)     1,696.8420 (2.29)     135.9121 (2.27)     1,644.7810 (2.29)     79.6613 (3.79)        36;37    589.3300 (0.44)        391           1
------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------

如何使 cythonized 版本更快?

【问题讨论】:

  • (这可能属于 codereview.stackexchange.com )没有实际查看您的代码,我知道这通常是 python 中的一个问题。您是否尝试过用堆栈替换递归调用?这将处理递归的逻辑,而无需实际进行函数调用。这是我在简短搜索中找到的现有最佳答案:stackoverflow.com/questions/13591970/…
  • 我假设您已经尝试生成 Cython 的带注释的 HTML 以查看它认为可能很慢的内容。我想知道递归函数是否有点像红鲱鱼,实际上不是你的问题。
  • @KennyOstrom 我不相信瓶颈是递归调用。 fibonacci 的 cython 版本比 python 版本快几个数量级blog.nelsonliu.me/2016/04/29/gsoc-week-0-cython-vs-python 考虑到它的递归很重,因此它或多或少也应该是正确的。
  • 我担心你在一个向量上调用len 会导致它被转换为一个列表然后len 被调用,例如。
  • 您是否尝试过使用编译器指令?在你的情况下,你正在做一个部门,我建议在你的班级定义之前添加 @cython.cdivision(False) 。在此处查看更多详细信息stackoverflow.com/questions/19537673/slow-division-in-cython

标签: python algorithm performance cython


【解决方案1】:

当尝试优化 cython 代码时,第一步是使用注释构建(例如参见 Cython-documentation 的这一部分),即

 cython -a xxx.pyx

或类似的。它会生成一个 html,在其中可以看到代码的哪些部分使用了 Python 功能。

在你看来,mid = (current_i + current_j)//2 是个问题。

它生成以下 C 代码:

  /*else*/ {
    __pyx_t_3 = __Pyx_PyInt_From_long(__Pyx_div_long((__pyx_v_current_i + __pyx_v_current_j), 2)); if (unlikely(!__pyx_t_3)) __PYX_ERR(0, 42, __pyx_L1_error)
    __Pyx_GOTREF(__pyx_t_3);
    __pyx_v_mid = __pyx_t_3;
    __pyx_t_3 = 0;

mid 是 Python-integer(由于 __Pyx_PyInt_From_long),所有使用它的操作都会导致更多地转换为 Python-integer 和缓慢的操作。

制作midcdef int。调查注释代码中的其他黄线(与 Python 的交互)。

【讨论】:

  • 感谢指点!我已经解决了这个问题,我注意到range_sumrange_minrange_max 函数是深黄色的,表明 python 交互很重。有什么办法可以解决这些问题吗?
  • 使它们成为声明返回类型 (int) 的 cpdef,与 combine-function 相同。 Cython 将选择 cdef 函数而不将结果转换为昂贵的 python 整数
  • 抱歉,我查错了版本。那就别担心:因为它们部分是 def 的,所以需要 python 交互。首先专注于代码。基准是什么样的?
  • 它们几乎相同,cython 大约快 1.18 倍
  • 得到一些改进!我没有将输入列表存储为向量,而是调整了代码,使其仅存储其长度。对于相同的基准,它的速度大约快 4.5 倍
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2021-11-27
  • 2012-04-25
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多