如何在Python完美平方检测代码中跳过特定(r,a,b)组合以提速
实现跳过两类(r,a,b)组合以提升完美平方查找代码效率
已知两类可直接跳过的正整数n对应的(r,a,b)组合:
- $r=3n+2, a=3n^2-2, b=1-3n+3n^3$
- $r=3n+17, a=73+30n+3n^2, b=361 + 222 n + 45 n^2 + 3 n^3$
现有Python代码运行速度极慢,已尝试实现第一类组合的跳过逻辑,但不知道如何处理第二类,需要完整实现跳过这两类组合的逻辑来提速。
用户现有判断逻辑片段:
def check_additional_conditions(r, a, b): """ Check if the additional conditions are satisfied """ # Condition 1: IntegerQ[(r-2)/3] == True condition1 = ((r - 2) / 3).is_integer # Condition 2: IntegerQ[Sqrt[(2 + a)/3]] == True condition2 = sp.sqrt((2 + a) / 3).is_integer # Condition 3: b = 1 - 3n + 3n^3 where n is a positive integer n = sp.Symbol('n', positive=True, integer=True) condition3 = sp.solve(1 - 3*n + 3*n**3 - b, n) condition3 = any([solution.is_integer and solution > 0 for solution in condition3]) return condition1 and condition2 and condition3
用户完整代码:
import sympy as sp import time import os from datetime import datetime from concurrent.futures import ProcessPoolExecutor def is_perfect_square(expr): """ Check if the expression is a perfect square """ sqrt_expr = sp.sqrt(expr) return sqrt_expr.is_integer or sqrt_expr == int(sqrt_expr) def calculate_b(r, a): """ Calculate the value of b based on the given formula """ expr = 3 * (3 * (-4 + r)**2 + 12 * a**2 * (-2 + r) - 4 * a * (-5 + r) * (-2 + r) + 4 * a**3 * (-2 + r)**2) b = (-12 + sp.sqrt(expr) + 3 * r) / (6 * (r - 2)) return b def monitor(r, a): """ Monitor function to print current values of r and a """ print(f"Checking r = {r}, a = {a}") def check_for_a_for_r(r, a_start=3): """ Main function to check conditions for a for a given r """ local_results = [] for a in range(a_start, 100000): expr = 3 * (3 * (-4 + r)**2 + 12 * a**2 * (-2 + r) - 4 * a * (-5 + r) * (-2 + r) + 4 * a**3 * (-2 + r)**2) monitor(r, a) if is_perfect_square(expr): b = calculate_b(r, a) if b.is_integer: local_results.append((r, a, b)) else: print(f"Skipping: r = {r}, a = {a}, b = {b} (not an integer)") return local_results def save_results(results): """ Save the results to a new text file with a unique filename """ script_directory = os.path.dirname(os.path.abspath(__file__)) timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") file_path = os.path.join(script_directory, f"perfect_squares_{timestamp}.txt") with open(file_path, "w") as file: if results: for r, a, b in results: file.write(f"r = {r}, a = {a}, b = {b}\n") print(f"Found perfect square: r = {r}, a = {a}, b = {b}") else: file.write("No perfect squares found.\n") results = [] last_save_time = time.time() with ProcessPoolExecutor() as executor: for r in range(3, 10**6 + 1): future = executor.submit(check_for_a_for_r, r, 3) # Start 'a' at 3 results.extend(future.result()) if time.time() - last_save_time >= 3600: save_results(results) last_save_time = time.time() save_results(results)
解决方案
核心优化思路
放弃用SymPy解方程的低效方式,直接通过r反推n,再验证a和b是否匹配对应公式,全程用整数运算避免精度和性能问题。
1. 编写高效的跳过判断函数
替换原有check_additional_conditions,实现同时检查两类组合的逻辑:
def should_skip(r, a, b): """判断当前(r,a,b)是否属于需要跳过的两类组合""" # 检查第一类组合: r=3n+2, a=3n²-2, b=1-3n+3n³ if (r - 2) % 3 == 0: n = (r - 2) // 3 if n > 0: expected_a1 = 3 * (n ** 2) - 2 expected_b1 = 1 - 3 * n + 3 * (n ** 3) if a == expected_a1 and b == expected_b1: return True # 检查第二类组合: r=3n+17, a=3n²+30n+73, b=3n³+45n²+222n+361 if (r - 17) % 3 == 0: n = (r - 17) // 3 if n > 0: expected_a2 = 3 * (n ** 2) + 30 * n + 73 expected_b2 = 3 * (n ** 3) + 45 * (n ** 2) + 222 * n + 361 if a == expected_a2 and b == expected_b2: return True return False
2. 在主逻辑中加入跳过判断
修改check_for_a_for_r函数,找到整数b后先判断是否需要跳过,再决定是否加入结果,同时移除低效的打印操作:
def check_for_a_for_r(r, a_start=3): """ Main function to check conditions for a for a given r """ local_results = [] for a in range(a_start, 100000): expr = 3 * (3 * (-4 + r)**2 + 12 * a**2 * (-2 + r) - 4 * a * (-5 + r) * (-2 + r) + 4 * a**3 * (-2 + r)**2) if is_perfect_square(expr): b = calculate_b(r, a) if b.is_integer: b_int = int(b) if not should_skip(r, a, b_int): local_results.append((r, a, b_int)) return local_results
3. 额外性能优化(可选)
优化完全平方数判断函数,用纯整数运算替代SymPy的符号计算:
def is_perfect_square(num): """检查整数是否为完全平方数""" if not isinstance(num, int): num = int(num) if num < 0: return False sqrt_num = int(sp.sqrt(num)) return sqrt_num * sqrt_num == num
完整修改后的代码
import sympy as sp import time import os from datetime import datetime from concurrent.futures import ProcessPoolExecutor def is_perfect_square(num): """检查整数是否为完全平方数""" if not isinstance(num, int): num = int(num) if num < 0: return False sqrt_num = int(sp.sqrt(num)) return sqrt_num * sqrt_num == num def calculate_b(r, a): """ Calculate the value of b based on the given formula """ expr = 3 * (3 * (-4 + r)**2 + 12 * a**2 * (-2 + r) - 4 * a * (-5 + r) * (-2 + r) + 4 * a**3 * (-2 + r)**2) b = (-12 + sp.sqrt(expr) + 3 * r) / (6 * (r - 2)) return b def should_skip(r, a, b): """判断当前(r,a,b)是否属于需要跳过的两类组合""" # 检查第一类组合: r=3n+2, a=3n²-2, b=1-3n+3n³ if (r - 2) % 3 == 0: n = (r - 2) // 3 if n > 0: expected_a1 = 3 * (n ** 2) - 2 expected_b1 = 1 - 3 * n + 3 * (n ** 3) if a == expected_a1 and b == expected_b1: return True # 检查第二类组合: r=3n+17, a=3n²+30n+73, b=3n³+45n²+222n+361 if (r - 17) % 3 == 0: n = (r - 17) // 3 if n > 0: expected_a2 = 3 * (n ** 2) + 30 * n + 73 expected_b2 = 3 * (n ** 3) + 45 * (n ** 2) + 222 * n + 361 if a == expected_a2 and b == expected_b2: return True return False def check_for_a_for_r(r, a_start=3): """ Main function to check conditions for a for a given r """ local_results = [] for a in range(a_start, 100000): expr = 3 * (3 * (-4 + r)**2 + 12 * a**2 * (-2 + r) - 4 * a * (-5 + r) * (-2 + r) + 4 * a**3 * (-2 + r)**2) if is_perfect_square(expr): b = calculate_b(r, a) if b.is_integer: b_int = int(b) if not should_skip(r, a, b_int): local_results.append((r, a, b_int)) return local_results def save_results(results): """ Save the results to a new text file with a unique filename """ script_directory = os.path.dirname(os.path.abspath(__file__)) timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") file_path = os.path.join(script_directory, f"perfect_squares_{timestamp}.txt") with open(file_path, "w") as file: if results: for r, a, b in results: file.write(f"r = {r}, a = {a}, b = {b}\n") print(f"Found perfect square: r = {r}, a = {a}, b = {b}") else: file.write("No perfect squares found.\n") results = [] last_save_time = time.time() with ProcessPoolExecutor() as executor: for r in range(3, 10**6 + 1): future = executor.submit(check_for_a_for_r, r, 3) # Start 'a' at 3 results.extend(future.result()) if time.time() - last_save_time >= 3600: save_results(results) last_save_time = time.time() save_results(results)
内容的提问来源于stack exchange,提问作者Jan Eerland
相关产品推荐
相关产品推荐

