如何优化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(给定数)以内的素数筛,分发到各节点后过滤倍数,再合并结果,这种思路能否提升性能?
优化方案
你的思路完全正确,这是分布式埃氏筛法的标准实现思路,能彻底解决当前的性能瓶颈,大幅提升处理大数值的效率。
核心优化逻辑
- 预先生成小素数集合:先生成
sqrt(n)以内的所有素数,这部分计算量很小,在Driver端即可快速完成。 - 广播小素数集合:将小素数集合通过Spark的广播变量分发到所有Executor节点,避免重复传输和计算。
- 分布式过滤:每个分区只需用广播的小素数集合过滤当前分区的数值,判断是否为素数。
这种方式下,所有分区的计算量基本一致,不会出现单个分区拖慢整体进度的情况,同时避免了每个分区重复计算小素数的冗余操作。
优化后的代码实现
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()
额外优化建议
- 减少待处理数据量:除了2以外,所有偶数都不是素数,因此可以只生成奇数+2的集合,直接减少一半的数据处理量,代码中已实现该优化。
- 分区大小调整:分区数量不宜过多或过少,建议每个分区的大小控制在10万-100万之间(根据集群资源调整),避免分区过小导致调度开销大,或分区过大导致单个任务耗时过长。
- 避免collect()大结果:当n达到10亿时,素数数量有约5000万,
collect()会将所有数据拉到Driver端,容易导致内存溢出,建议直接用saveAsTextFile()保存结果到分布式存储(如HDFS)。
内容的提问来源于stack exchange,提问作者Lemon
相关产品推荐
相关产品推荐

