代码之家  ›  专栏  ›  技术社区  ›  Ξένη Γήινος

如何优化Eratosthenes的NumPy筛?

  •  0
  • Ξένη Γήινος  · 技术社区  · 3 年前

    我已经在NumPy中实现了我自己的埃拉托斯梯尼筛。我相信你们都知道这是为了找到一个数下的所有素数,所以我不做任何进一步的解释。

    代码:

    import numpy as np
    
    def primes_sieve(n):
        primes = np.ones(n+1, dtype=bool)
        primes[:2] = False
        primes[4::2] = False
        for i in range(3, int(n**0.5)+1, 2):
            if primes[i]:
                primes[i*i::i] = False
    
        return np.where(primes)[0]
    

    正如你所看到的,我已经做了一些优化,首先,除了2之外,所有素数都是奇数,所以我将2的所有倍数设置为 False 并且只有蛮力奇数。

    其次,我只循环遍历数字,直到平方根的底部,因为平方根之后的所有复数都会被平方根以下素数的倍数所消除。

    但它不是最优的,因为它循环通过低于极限的所有奇数,并且不是所有奇数都是素数。随着数量的增加,素数变得越来越稀疏,所以有很多冗余的迭代。

    因此,如果候选列表是动态变化的,以这种方式,已经识别的复合数甚至永远不会被迭代,因此只有素数循环通过,就不会有任何浪费的迭代,因此算法将是最优的。

    我写了一个优化版本的粗略实现:

    def primes_sieve_opt(n):
        primes = np.ones(n+1, dtype=bool)
        primes[:2] = False
        primes[4::2] = False
        limit = int(n**0.5)+1
        i = 2
        while i < limit:
            primes[i*i::i] = False
            i += 1 + primes[i+1:].argmax()
    
        return np.where(primes)[0]
    

    但它比未优化的版本慢得多:

    In [92]: %timeit primes_sieve(65536)
    271 µs ± 22 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
    
    In [102]: %timeit primes_sieve_opt(65536)
    309 µs ± 3.86 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
    

    我的想法很简单,通过获取的下一个索引 True ,我可以确保所有素数都被覆盖,并且只处理素数。

    然而 np.argmax 在这方面进展缓慢。我在谷歌上搜索了“如何在NumPy数组中找到下一个True值的索引”(没有引号),我发现了几个StackOverflow问题,这些问题稍微相关,但最终没有回答我的问题。

    例如 numpy get index where value is true Numpy first occurrence of value greater than existing value

    我并没有试图找到所有索引 真的 ,这样做是非常愚蠢的,我需要找到下一个 真的 值,获取其索引并立即停止循环,只有 bool s

    我如何优化它?


    编辑

    如果有人感兴趣,我已经进一步优化了我的算法:

    import numba
    import numpy as np
    
    @numba.jit(nopython=True, parallel=True, fastmath=True, forceobj=False)
    def prime_sieve(n: int) -> np.ndarray:
        primes = np.full(n + 1, True)
        primes[:2] = False
        primes[4::2] = False
        primes[9::6] = False
        limit = int(n**0.5) + 1
        for i in range(5, limit, 6):
            if primes[i]:
                primes[i * i :: 2 * i] = False
    
        for i in range(7, limit, 6):
            if primes[i]:
                primes[i * i :: 2 * i] = False
    
        return np.flatnonzero(primes)
    

    我用过 numba 加快速度。由于除了2和3之外的所有素数都是6k+1或6k-1,这使得事情变得更快。

    1 回复  |  直到 3 年前
        1
  •  5
  •   Nick ODell    3 年前

    我的想法很简单,通过获得True的下一个索引,我可以确保所有素数都被覆盖,并且只处理素数。

    一些分析表明,通过这种方式,你最多可以获得0.2%的加速。

    对于N的大值,绝大多数时间都花在 primes[i*i::i] = False

    以下是在前一亿个素数上运行的line_profiler的输出:

    Timer unit: 1e-09 s
    
    Total time: 1.04878 s
    File: /tmp/ipykernel_22262/2557137730.py
    Function: primes_sieve at line 3
    
    Line #      Hits         Time  Per Hit   % Time  Line Contents
    ==============================================================
         3                                           def primes_sieve(n):
         4         1   14264754.0 14264754.0      1.4      primes = np.ones(n+1, dtype=bool)
         5         1      12394.0  12394.0      0.0      primes[:2] = False
         6         1   16238905.0 16238905.0      1.5      primes[4::2] = False
         7      4999    1309955.0    262.0      0.1      for i in range(3, int(n**0.5)+1, 2):
         8      3771    1507909.0    399.9      0.1          if primes[i]:
         9      1228  909007228.0 740233.9     86.7              primes[i*i::i] = False
        10                                           
        11         1  106434647.0 106434647.0     10.1      return np.where(primes)[0]
    

    如果您跳过了的更多值 i ,你可以避免花在排队上的时间 for i in range(3, int(n**0.5)+1, 2): if primes[i]: 。但你无法避免在 素数[i*i::i]=假 。由于程序在每一个方面都花费了0.1%,因此最多可以节省0.2%的执行时间。