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

如何优化PySpark版埃氏筛法以支持十亿级大数范围的可扩展性?

问题描述

尝试编写PySpark脚本生成≤给定数值的所有素数(例如10亿以内),现有代码在小数值下表现良好,但数值达到1亿后性能急剧下降。

现有代码如下:

from pyspark.sql import SparkSession
from math import isqrt

def is_perfect_square(num):
  root = isqrt(num)
  return root*root == num


def sieve_of_eratosthenes_partition(iterator):
    upper_limit = max(iterator) # Upper limit for prime number generation
    prime_flag = [True] * len(iterator) # Initialize boolean array for primes
    result = []
    cur_prime = 2

    while cur_prime * cur_prime <= upper_limit:
        i = 0
        if ((cur_prime % 2 == 0 and cur_prime != 2) or is_perfect_square(cur_prime)):
          cur_prime += 1
        else:
          for num in iterator:
            if(num % cur_prime == 0 and num != cur_prime):
              prime_flag[i] = False
            i += 1  
          cur_prime += 1

    for num, is_prime in zip(iterator, prime_flag):
        if is_prime and num > 1:
            result.append(num)
    return result

spark = SparkSession.builder.appName("SievePrimesMapPartitions").getOrCreate()
n = 10**7 # End range 

numbers = spark.sparkContext.parallelize(range(1, n+1), 1000)
result_rdd = numbers.mapPartitions(sieve_of_eratosthenes_partition)

# result_rdd.map(str).saveAsTextFile("primes")

primes = result_rdd.collect()
print(primes)
print(len(primes))

当前实现思路:将1到1000万划分为1000个均匀分区,对每个分区单独应用筛法,迭代过滤素数的倍数。

性能瓶颈:数值大的分区(如9990000…10000000)需要循环到sqrt(1000万)才能完成筛选,而小分区只需循环到100,整体性能由最慢的分区决定。

疑问:

  • 如何改进现有代码?
  • 是否有更合理的分区方式?
  • 先生成sqrt(给定数)以内的素数筛,分发到各节点后过滤倍数,再合并结果,这种思路能否提升性能?

优化方案

你的思路完全正确,这是分布式埃氏筛法的标准实现思路,能彻底解决当前的性能瓶颈,大幅提升处理大数值的效率。

核心优化逻辑

  1. 预先生成小素数集合:先生成sqrt(n)以内的所有素数,这部分计算量很小,在Driver端即可快速完成。
  2. 广播小素数集合:将小素数集合通过Spark的广播变量分发到所有Executor节点,避免重复传输和计算。
  3. 分布式过滤:每个分区只需用广播的小素数集合过滤当前分区的数值,判断是否为素数。

这种方式下,所有分区的计算量基本一致,不会出现单个分区拖慢整体进度的情况,同时避免了每个分区重复计算小素数的冗余操作。

优化后的代码实现

from pyspark.sql import SparkSession
from math import isqrt

def sieve_small_primes(n):
    """生成sqrt(n)以内的所有素数,用于后续过滤"""
    if n < 2:
        return []
    sieve = [True] * (n + 1)
    sieve[0] = sieve[1] = False
    for i in range(2, isqrt(n) + 1):
        if sieve[i]:
            sieve[i*i : n+1 : i] = [False] * len(sieve[i*i : n+1 : i])
    return [i for i, is_prime in enumerate(sieve) if is_prime]

def filter_primes_partition(iterator, small_primes):
    """用预先生成的小素数过滤当前分区的数值"""
    primes = []
    for num in iterator:
        if num < 2:
            continue
        # 检查是否能被小素数整除
        is_prime = True
        for p in small_primes:
            if p * p > num:
                break
            if num % p == 0:
                is_prime = False
                break
        if is_prime:
            primes.append(num)
    return primes

if __name__ == "__main__":
    spark = SparkSession.builder.appName("DistributedSieve").getOrCreate()
    n = 10**7  # 目标数值范围,可扩展到10^9
    
    # 步骤1:生成sqrt(n)以内的小素数
    sqrt_n = isqrt(n)
    small_primes = sieve_small_primes(sqrt_n)
    
    # 步骤2:广播小素数集合到所有节点
    broadcast_small_primes = spark.sparkContext.broadcast(small_primes)
    
    # 步骤3:生成待处理的数值RDD,可优化为只生成奇数(除了2)减少数据量
    # 优化点:排除偶数,只处理奇数+2,减少一半数据量
    numbers_rdd = spark.sparkContext.parallelize([2] + list(range(3, n+1, 2)), 1000)
    
    # 步骤4:每个分区用广播的小素数过滤
    result_rdd = numbers_rdd.mapPartitions(
        lambda iter: filter_primes_partition(iter, broadcast_small_primes.value)
    )
    
    # 保存或收集结果(大数值场景建议直接保存,避免collect()内存溢出)
    # result_rdd.map(str).saveAsTextFile("primes")
    primes = result_rdd.collect()
    print(f"找到的素数数量:{len(primes)}")
    
    spark.stop()

额外优化建议

  1. 减少待处理数据量:除了2以外,所有偶数都不是素数,因此可以只生成奇数+2的集合,直接减少一半的数据处理量,代码中已实现该优化。
  2. 分区大小调整:分区数量不宜过多或过少,建议每个分区的大小控制在10万-100万之间(根据集群资源调整),避免分区过小导致调度开销大,或分区过大导致单个任务耗时过长。
  3. 避免collect()大结果:当n达到10亿时,素数数量有约5000万,collect()会将所有数据拉到Driver端,容易导致内存溢出,建议直接用saveAsTextFile()保存结果到分布式存储(如HDFS)。

内容的提问来源于stack exchange,提问作者Lemon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 23:55:25