如何高效生成正整数N的无重复两因子分解对
问题
给定正整数N,需要找到所有满足N = a*b的正整数对(a, b),且要求避免(a,b)与(b,a)这类重复结果(即a和b视为同一组分解)。
目前使用sympy.ntheory.factorint进行质因数分解,该函数返回以质因数为键、指数为值的字典。现有代码如下:
import itertools import numpy as np from sympy.ntheory import factorint def find_decompositions(n): prime_factors = factorint(n) cut_points = {f: [i for i in range(1+e)] for f, e in prime_factors.items()} cuts = itertools.product(*cut_points.values()) decompositions = [((a := np.prod([f**e for f, e in zip(prime_factors, cut)])), n//a) for cut in cuts] return decompositions
示例运行结果:
In [235]: find_decompositions(12) Out[235]: [(1, 12), (3, 4), (2, 6), (6, 2), (4, 3), (12, 1)]
期望得到的结果:
Out[235]: [(1, 12), (3, 4), (2, 6)]
尝试过调整cut_points中的范围(如e//2、1 + e//2等)但未成功,不想用decompositions[:(len(decompositions)+1)//2]这种截取前半部分的方法,希望找到直接减少计算量的解决方案。
解决方案
要从根源减少计算量,我们只需要生成满足a ≤ b的分解对——也就是确保计算出的a不超过sqrt(N),这样就能避免生成重复的(b,a)项。
递归优化版(高效剪枝)
import itertools import numpy as np from sympy.ntheory import factorint def find_decompositions(n): prime_factors = factorint(n) primes = list(prime_factors.keys()) exponents = list(prime_factors.values()) sqrt_n = np.sqrt(n) def generate_valid_a(idx, current_product): if idx == len(primes): if current_product <= sqrt_n: yield current_product return max_k = exponents[idx] for k in range(max_k + 1): next_product = current_product * (primes[idx] ** k) if next_product > sqrt_n: break # 超过平方根后,更大的指数只会让乘积更大,直接终止该分支 yield from generate_valid_a(idx + 1, next_product) return [(a, n // a) for a in generate_valid_a(0, 1)]
代码说明
- 递归遍历每个质因数的指数组合,实时计算当前
a的乘积 - 一旦乘积超过
sqrt(N),立即终止该分支的遍历,避免无效计算 - 仅保留
a ≤ sqrt(N)的组合,从根源上消除重复项,计算量直接减半甚至更少(因剪枝提前终止)
测试结果:
In [236]: find_decompositions(12) Out[236]: [(1, 12), (2, 6), (3, 4)]
迭代过滤版(简单易读)
如果觉得递归不好理解,也可以在遍历所有指数组合时直接过滤掉a > sqrt(N)的项,相比原代码仍能减少一半的无效结果处理:
import itertools import numpy as np from sympy.ntheory import factorint def find_decompositions(n): prime_factors = factorint(n) primes = list(prime_factors.keys()) exponent_ranges = [range(e + 1) for e in prime_factors.values()] sqrt_n = np.sqrt(n) decompositions = [] for cuts in itertools.product(*exponent_ranges): a = np.prod([p**k for p, k in zip(primes, cuts)]) if a > sqrt_n: continue decompositions.append((a, n // a)) return decompositions
内容的提问来源于stack exchange,提问作者Nick Skywalker
相关产品推荐
相关产品推荐

