论文复现:多项式化简实现性能优化求助
SymPy多项式降阶的性能优化方案
问题背景
我正在复现某论文的多项式降阶逻辑,小参数场景下(如N=15,l1=2,l2=3)能得到与论文一致的结果,但当参数增大到N=10403、l1=l2=7时,生成最终表达式耗时长达数分钟,急需优化性能。
原实现代码:
from sympy import IndexedBase, expand, Indexed, Mul import numpy as np from tqdm import tqdm x = IndexedBase('x') def high_degree_f(N, l1 = 2, l2 = 3): # p and q are binary numbers of length l1 and l2 p = 1 q = 1 for i in range(1, l1): p += x[i] * 2**i for idx, l in enumerate(range(l1, l1+l2 - 1)): q += x[l] * 2**(idx+1) f = (N - p * q ) **2 f = expand(f) r = len(f.free_symbols) + 1 # replace x[i]**k by x[i] as x[i] is binary for i in range(1, r): for k in range(2, 4): f = f.subs(x[i]**k, x[i]) return (f, p, q) def out_degrees_more_than_2(f): # Iterate over the terms in the expanded polynomial for term in f.as_ordered_terms(): # Check if the term is a product if isinstance(term, Mul): # Count the number of Indexed instances in the term variable_count = sum(isinstance(factor, Indexed) for factor in term.args) # Check if there are more than two variables in the product if variable_count > 2: # Print the term yield (variable_count, term) def max_degree(f): max_degree = 0 for (variable_count, term) in out_degrees_more_than_2(f): max_degree = max(max_degree, variable_count) return max_degree def reduced_f(N, l1 = 2, l2 = 3): (f, p, q) = high_degree_f(N, l1 = l1, l2 = l2) number_of_variables = len(f.free_symbols) - 1 while max_degree(f) > 2: for (variable_count, term) in tqdm(out_degrees_more_than_2(f)): if variable_count == 3: # extract the numerical coefficient of the term: coefficient, v = term.as_coeff_Mul() variables = v.args new_term = coefficient* (variables[2] * x[number_of_variables+1] + 2 * ( variables[0] * variables[1] - 2 * variables[0] * x[number_of_variables+1] - 2 * variables[1] * x[number_of_variables+1] + 3 * x[number_of_variables+1])) number_of_variables += 1 # substitute the term with the new term f = f.subs(term, new_term) elif variable_count == 4: pass else: raise ValueError("Unexpected number of variables in term") f = expand(f) return (f, p, q) N = 15 l1 = 2 l2 = 3 (f, p, q) = reduced_f(N, l1 = l1, l2 = l2) print("N:\t\t", N) print("p:\t\t", p) print("q:\t\t", q) print("Reduced f:\t", f)
小参数运行结果:
N: 15 p: 2*x[1] + 1 q: 2*x[2] + 4*x[3] + 1 Reduced f: 200*x[1]*x[2] - 48*x[1]*x[3] - 512*x[1]*x[4] - 52*x[1] + 16*x[2]*x[3] - 512*x[2]*x[4] - 52*x[2] + 128*x[3]*x[4] - 96*x[3] + 768*x[4] + 196
论文对应表达式:
$$\begin{array}{rcl}f^{\prime} (x) & = & 200{x}{1}{x}{2}-48{x}{1}{x}{3}-512{x}{1}{x}{4}+16{x}{2}{x}{3}-512{x}{2}{x}{4}+128{x}{3}{x}{4}\ & & -52{x}{1}-52{x}{2}-96{x}{3}+768{x}{4}+\mathrm{196,}\end{array}$$
核心优化点及代码修改
1. 避免反复展开与重复遍历
原代码每次替换单个三次项后反复调用expand,且多次遍历多项式计算最大次数,这是主要性能瓶颈。改为批量收集所有三次项,一次性生成替换规则,完成替换后再展开,大幅减少遍历和展开次数。
2. 简化替换表达式并预计算变量数量
- 提前计算初始变量数量(无需通过
free_symbols查询):初始变量数为l1 + l2 - 2(p用l1-1个变量,q用l2-1个)。 - 化简三次项的替换公式,合并同类项,降低后续展开的复杂度。
3. 优化最大次数检查逻辑
不再每次遍历所有项计算max_degree,而是在批量替换三次项后,直接检查是否还有剩余的三次项,若无则退出循环。
优化后的完整代码
from sympy import IndexedBase, expand, Indexed, Mul import numpy as np from tqdm import tqdm x = IndexedBase('x') def high_degree_f(N, l1=2, l2=3): p = 1 q = 1 for i in range(1, l1): p += x[i] * 2**i for idx, l in enumerate(range(l1, l1 + l2 - 1)): q += x[l] * 2**(idx + 1) f = (N - p * q)**2 f = expand(f) # 直接计算初始变量数,无需查询free_symbols num_vars = l1 + l2 - 2 # 替换x[i]^k为x[i](二进制变量特性) for i in range(1, num_vars + 1): f = f.subs(x[i]**2, x[i]) f = f.subs(x[i]**3, x[i]) return f, p, q, num_vars def collect_cubic_terms(f): """批量收集所有三次项""" cubic_terms = [] for term in f.as_ordered_terms(): if isinstance(term, Mul): var_count = sum(isinstance(factor, Indexed) for factor in term.args) if var_count == 3: cubic_terms.append(term) return cubic_terms def reduced_f(N, l1=2, l2=3): f, p, q, num_vars = high_degree_f(N, l1=l1, l2=l2) while True: cubic_terms = collect_cubic_terms(f) if not cubic_terms: break # 无三次项,退出循环 # 批量生成替换规则 subs_map = {} for term in tqdm(cubic_terms, desc="处理三次项"): coeff, vars_mul = term.as_coeff_Mul() vars_list = list(vars_mul.args) new_idx = num_vars + 1 # 化简替换表达式:合并同类项 new_term = coeff * ( 2 * vars_list[0] * vars_list[1] + x[new_idx] * (vars_list[2] - 4*vars_list[0] -4*vars_list[1] +6) ) subs_map[term] = new_term num_vars += 1 # 一次性替换所有三次项,再展开 f = f.subs(subs_map) f = expand(f) # 替换新变量的高次幂(新变量也是二进制) for i in range(num_vars - len(cubic_terms) + 1, num_vars + 1): f = f.subs(x[i]**2, x[i]) f = f.subs(x[i]**3, x[i]) return f, p, q # 测试小参数 N = 15 l1 = 2 l2 = 3 f, p, q = reduced_f(N, l1=l1, l2=l2) print("N:\t\t", N) print("p:\t\t", p) print("q:\t\t", q) print("Reduced f:\t", f) # 大参数测试(取消注释即可运行) # N = 10403 # l1 = 7 # l2 = 7 # f, p, q = reduced_f(N, l1=l1, l2=l2)
额外优化建议
如果上述优化后速度仍不满足需求,可以考虑避开SymPy的符号计算开销:用字典存储多项式的项(键为变量索引的元组,值为系数),手动实现项的遍历、替换和合并逻辑,这种方式的计算速度会远快于SymPy的符号操作,尤其适合大规模多项式处理。
内容的提问来源于stack exchange,提问作者Amirhossein Rezaei
相关产品推荐
相关产品推荐

