【问题标题】:Reducing the time complexity of stair-step question (Amazon interview question)降低阶梯式问题的时间复杂度(亚马逊面试问题)
【发布时间】:2023-03-30 19:05:01
【问题描述】:
def step(n):
    if (n==0) or (n==1):
        return 1 
    elif n==2:
        return 2
    else:
        return step(n-1) + step(n-2) + step(n-3)
n = int(input())
print(step(n))

输入 53798080 需要 1 秒。满足测试用例的时间应该比这要少得多。

【问题讨论】:

  • 听起来更像是一道数学题。还会使用step(-1) 或任何浮点数引发堆栈溢出。
  • 不提供浮点输入。
  • 可能首先注意到您多次计算相同的步骤。
  • 你帮我指出来可以吗?
  • 函数式语言通常通过缓存结果来加速纯函数。你可以这样做。我用一个循环解决了这个问题(因此 O(N) ),但你不能使用递归来做到这一点。

标签: python recursion


【解决方案1】:

这类问题 - 评估递归关系 - 多年来已经有很多聪明人研究它,这意味着您可以使用大量很酷的见解和想法来加快速度。

cmets 已经很好地确定了为什么您的代码在大输入时会变慢 - 这是因为您正在生成大量重复的递归调用。那么,问题是如何解决这个问题。

如果您想保持相同的基本策略,我建议您使用memoization。如果您以前没有见过这种技术,基本思想是让递归跟踪已经进行的调用并缓存这些调用的结果。然后,如果您尝试两次解决相同的问题,您可以将缓存的结果交回。

记忆的通用模板看起来像这样。 (它是伪代码,但不应该太难适应。)

def memoized_recursion(original_args, memoization_table):
    if memoization_table contains original_args):
        return memoization_table[original_args]
    else
        # Put the rest of your recursive code here.
        # Before returning a result, store it in memoization_table.

这极大地减少了递归调用的数量,从而加快了您的代码速度。

当然,这并不是使您的代码更快的唯一解决方案。如果您必须保持递归,则可以使用不同的洞察力从根本上改变策略。基本思路是这样的。您正在生成一系列如下所示的数字:

1, 1, 2, 4, 7, 13, 24, ...

想法是这样的

  • 前三项分别为 1、1、2;
  • 这个之后的每一项都是前面三个数字的总和;和
  • 您想要该系列的第 n 个学期。

如果您需要术语 0、1 或 2,您可以直接阅读答案,因为您知道前三个数字。

如果没有,您可以使用另一种技术。与其获取前面的三个值并将它们相加,不如使用这个有用的事实:要求以 1、1、2 开头的系列的第 n 项等同于要求以 1、1、2 开头的系列的第 (n-1) 项1、2、4。(你明白为什么吗?)

更一般地,如果系列的前三个项是 a、b 和 c,并且您想要第 n 个项,您可以要求从序列 b、c 开始的系列的第 (n-1) 个项, a + b + c。这提供了一种不同的递归策略,其中递归不分支,这意味着您不需要记忆。

现在,最后一个策略。您要解决的问题类型涉及一种称为齐次线性递推关系的东西。也就是说,你有一个重复的形式

  • a0、a1、...、ak-1是固定常数,
  • an+k = c0 an + c1 an+1 + ... + ck-1an+k-1.

这种重复包括斐波那契数列、佩尔数、帕多万数列等。

事实证明,在任何情况下,如果您要解决这样的递归,您都可以通过将特定选择的矩阵提高到特定的幂来解决问题。在您的情况下,基本思想与第二种递归策略的思想有关。这个想法是,如果序列的最后三个项是 a、b 和 c,那么您知道下一项是 a + b + c,而在这之前的两个项是 b 和 c。换句话说,您可以想象一个将 (a, b, c) 转换为 (b, c, a + b + c) 的映射。这可以被认为是这个矩阵方程:

| 0  1  0 | |a|   |     b     |
| 0  0  1 | |b| = |     c     |
| 1  1  1 | |c|   | a + b + c |

如果你让 M 是最左边的矩阵,那么计算 Mn 并将其乘以列向量 (a, b, c) 将得到第 n、(n+1)st 和 (n+2) )nd 递归关系的项。这给出了解决问题的完全不同的策略:构建一个矩阵,然后将其提升到一个大幂!

事实上,您可以非常有效地做到这一点。有一种称为exponentiation by squaring 的(递归)技术可以仅使用 O(log n) 乘法来计算矩阵的 n 次方。 (不幸的是,矩阵的条目将开始变得非常大,并且将它们相乘将开始成为您的瓶颈)。不过,这个策略可能值得一试,因为它是一种非常酷的技术!

最后,还有最后一个选项。如果您进行一些谷歌搜索,您会发现您的问题与找到第 n 个tribonacci number 密切相关。您可以使用一些很酷的公式来直接计算它,也涉及数字的幂,尽管它们可能会引入一些舍入错误,从而对您的目的来说太慢了。

【讨论】:

    【解决方案2】:

    输入 53798080 需要 1 秒。

    我非常怀疑这一点。您的代码堆栈在此输入上溢出。我认为这里发生的情况是输入 30 需要 1 秒才能产生 输出 53798080。通过输入 31,我们达到了将近 一半一分钟

    如果我们记住您的代码:

    from functools import lru_cache
    
    @lru_cache
    def step(n):
        if n == 0 or n == 1:
            return 1
    
        if n == 2:
            return 2
    
        return step(n-1) + step(n-2) + step(n-3)
    

    正如@templatetypedef 解释的那样,它解决了速度问题。但是,它会在输入 500 以上时因堆栈溢出(假设您没有分配更多堆栈)而爆炸。我们可以将该范围加倍,并使用更有效且递归更少的算法来处理速度问题而无需记忆:

    def step(n, prev1=2, prev2=1, prev3=1):
        if 0 <= n <= 1:
            return 1
    
        if n == 2:
            return prev1
    
        return step(n - 1, prev3 + prev2 + prev1, prev1, prev2)
    

    这将处理高达 999 的输入并在几分之一秒内产生结果:

    > time python3 test.py
    1499952522327196729941271196334368245775697491582778125787566254148069690528296568742385996324542810615783529390195412125034236407070760756549390960727215226685972723347839892057807887049540341540394345570010550821354375819311674972209464069786275283520364029575324
    0.032u 0.011s 0:00.04 100.0%    0+0k 0+0io 0pf+0w
    >
    

    (在此代码中添加@lru_cache 会将输入范围缩小回原来的范围,并且在速度方面没有任何区别。)

    【讨论】:

      猜你喜欢
      • 2011-08-30
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2022-06-28
      • 1970-01-01
      • 2015-06-15
      • 2015-10-31
      相关资源
      最近更新 更多