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

论文复现:多项式化简实现性能优化求助

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 18:35:54