跳转至

第 24 章 多项式

配套例题:BISHI18 多项式输出 来源:S3 day9《在数学花园里闲逛》的「多项式及其性质」「霍纳法则」「多项式差分」

多项式在笔试里主要以两种形态出现:格式化输出(考的是分支的严密性)和 求值/系数运算(考的是算法)。本章两者都讲。


24.1 表示

一元 \(n\) 次多项式

\[f(x) = a_n x^n + a_{n-1} x^{n-1} + \cdots + a_1 x + a_0\]

在代码里就是一个系数数组。有两种存法,各有适用场景:

存法 形式 优点
低位在前 a[i]\(x^i\) 的系数 下标即次数,乘法/求值最自然
高位在前 a[0]\(x^n\) 的系数 和输入顺序一致,输出最自然

BISHI18 的输入是高位在前(先给 \(a_n\)),而多数算法(乘法、求值)用低位在前更顺。 读入后 a.reverse() 一次即可,别在两种约定之间反复横跳——这是这类题最常见的 bug 来源。


24.2 求值:霍纳法则(秦九韶算法)

朴素求值要算 \(x^k\),是 \(O(n^2)\)\(O(n \log n)\)(配快速幂)。 霍纳法则把多项式改写成嵌套形式:

\[f(x) = a_0 + x\big(a_1 + x(a_2 + \cdots + x(a_{n-1} + x \cdot a_n)\cdots)\big)\]

只需要 \(n\) 次乘法和 \(n\) 次加法,\(O(n)\),且不产生任何幂次的中间大数

def horner(a, x):
    """a 为低位在前的系数列表,返回 f(x)。O(n)。"""
    res = 0                        # 初值 0 表示「还没吃进任何一项」
    for c in reversed(a):          # 从最高次开始,每轮把已算部分乘 x 再补一个系数
        res = res * x + c          # 全程只有乘和加,不会出现 x 的幂这种中间大数
    return res

取模版本(竞赛里更常用):

def horner_mod(a, x, mod):
    res = 0
    for c in reversed(a):
        res = (res * x + c) % mod  # 每步取模,res 始终小于 mod,乘法不会退化成大整数
    return res

为什么霍纳法则在 Python 里格外重要:不用它的话, \(x^n\)\(x\) 稍大时会变成巨大的整数(比如 \(x=10, n=100\) 就是 100 位数), 每次乘法都退化成大整数运算。霍纳法则配合取模,全程都是小整数。 这和 22-高精度与大整数 §22.2 的准则一是同一回事。


24.3 多项式的加法与乘法

def poly_add(a, b):
    """低位在前。O(max(n, m))。"""
    n = max(len(a), len(b))                # 和的次数是两者次数的较大者
    # 短的那个多项式在高次项上补 0,两边同次的系数直接相加
    return [(a[i] if i < len(a) else 0) + (b[i] if i < len(b) else 0)
            for i in range(n)]


def poly_mul(a, b):
    """朴素卷积。O(n*m)。"""
    res = [0] * (len(a) + len(b) - 1)      # n 次乘 m 次得 n+m 次,共 n+m+1 项
    for i, x in enumerate(a):
        if x == 0:
            continue                       # 系数为 0,整轮内层都是加 0,跳过
        for j, y in enumerate(b):
            res[i + j] += x * y            # x^i 乘 x^j 落在 x^(i+j) 上,下标直接相加
    return res

\(n, m\)\(10^4\)\(O(nm) = 10^8\) 次纯 Python 循环必然 TLE。 更大的规模需要 FFT/NTT——但这已超出笔试范围(S3 课件里也标注为省选内容), 本教程不展开,00-知识大纲 里已把 FFT 列为明确排除项。

一个实用的加速:如果只是求两个数列的卷积且系数不大, 可以把系数打包成大整数做一次乘法(Kronecker 替换)—— Python 的大整数乘法是 C 实现的 Karatsuba,往往比 Python 层双重循环快得多:

def poly_mul_fast(a, b, bits=64):
    """把系数打包进一个大整数做乘法。要求系数非负且卷积结果 < 2^bits。"""
    # 把每个系数塞进一个 bits 位宽的「格子」:整数的第 i 个格子 = x^i 的系数
    A = sum(c << (bits * i) for i, c in enumerate(a))
    B = sum(c << (bits * i) for i, c in enumerate(b))
    C = A * B                              # 大整数乘法本身就在算卷积,且走 C 层
    mask = (1 << bits) - 1                 # 低 bits 位全 1,用来抠出单个格子
    n = len(a) + len(b) - 1
    # 右移 bits*i 把第 i 个格子挪到最低位,再与 mask 取与
    return [(C >> (bits * i)) & mask for i in range(n)]

用之前必须确认 bits 足够大,否则相邻系数会互相污染。


24.4 差分找规律

S3 课件《在数学花园里闲逛》里有个很实用的结论:

\(n\) 次多项式的数列,做 \(n\) 次差分后变成常数列。

这给了一个识别「数列是不是多项式」的方法,也给了一个根据前几项猜通项的手段:

def diff_table(a):
    """打印差分表,直到出现常数行。"""
    cur = a[:]                             # 复制一份,不动调用方传进来的数列
    # 停止条件:只剩一项,或者整行元素都相同(set 去重后只剩一个值)
    while len(cur) > 1 and len(set(cur)) > 1:
        print(cur)
        # 相邻两项相减,每做一次差分数列就短一项、次数就降一次
        cur = [cur[i + 1] - cur[i] for i in range(len(cur) - 1)]
    print(cur)                             # 最后这一行就是常数行

例如数列 1, 8, 27, 64, 125(即 \(n^3\)):

[1, 8, 27, 64, 125]
[7, 19, 37, 61]
[12, 18, 24]
[6, 6]

3 次差分后变常数 6,说明是 3 次多项式,且最高次系数是 \(6/3! = 1\)

笔试里的用法:打表算出前几项,做差分,若很快变常数就说明有多项式通项, 可以直接拉格朗日插值或解方程求出来,避免推导。差分本身见 42-前缀和与差分


24.5 例题:BISHI18 多项式输出(简单)

给定 \(n\) 次多项式的系数(高位在前,\(-100 \le a_i \le 100\)\(a_n \ne 0\)), 按规则格式化输出:

  • 从高次到低次;
  • 系数为 0 的项完全省略
  • 次数 \(\ge 1\) 且系数为 \(\pm 1\) 时,省略系数的绝对值 1(常数项即使是 \(\pm 1\) 也要完整输出);
  • 次数 0 只输出常数;次数 1 输出 x;次数 \(\ge 2\) 输出 x^k
  • 第一个非零项若为正不输出前导 +,后续正项加 +,负项加 -

样例:5 / 100 -1 1 -3 0 10100x^5-x^4+x^3-3x^2+10

这题没有任何算法,全是分支的严密性\(n \le 100\),规模毫无压力。 正确的写法是把每一项的输出拆成三段独立决策,不要写成一个巨大的 if-else 链:

  1. 符号段:是第一项且为正 → 空;正 → +;负 → -
  2. 系数段:次数为 0 → 输出 abs(a)abs(a) == 1 → 空;否则 → abs(a)
  3. 变量段:次数 0 → 空;次数 1 → x;否则 → x^k
def main():
    n = int(input())
    a = list(map(int, input().split()))          # 高位在前:a[0] 是 x^n 的系数

    parts = []                                    # 每个非零项拼成一段,最后一次性接起来
    for i, c in enumerate(a):
        k = n - i                                 # 当前次数:下标每前进 1,次数就降 1
        if c == 0:                                # 规则:系数为 0 完全省略
            continue

        # 1) 符号
        if not parts:                             # parts 还空 ⇒ 这是第一个非零项
            sign = "-" if c < 0 else ""           # 首项为正时不输出前导加号
        else:
            sign = "-" if c < 0 else "+"          # 非首项一律带符号,正号也要写

        # 2) 系数:次数 >= 1 且绝对值为 1 时省略
        v = abs(c)                                # 符号已经单独处理,这里只要绝对值
        coef = "" if (k >= 1 and v == 1) else str(v)   # k >= 1 保住了常数项的 1

        # 3) 变量
        if k == 0:
            var = ""                              # 常数项没有 x
        elif k == 1:
            var = "x"                             # 一次项写 x,不写 x^1
        else:
            var = "x^{}".format(k)

        parts.append(sign + coef + var)           # 三段拼成完整一项

    print("".join(parts))                         # 中间不加任何分隔符,符号已在段内


main()

三个容易写错的地方:

  1. 常数项的 \(\pm 1\) 不能省略x^2+1 里的 1 必须留着, 否则会输出成 x^2+。所以省略条件必须带上 k >= 1
  2. 「第一项」指的是第一个非零项,不是 \(i = 0\)。 题目保证 \(a_n \ne 0\) 所以这里恰好一致,但把判断写成 if not parts 才是对任意输入都正确的写法。
  3. 符号要和系数分离。写成 str(c) 会得到 -3, 再拼上前面的 + 就成了 +-3。必须先取 abs,符号单独处理。

完整题解见 solutions/BISHI18.py,已通过官方样例验证。


24.6 本章速查

场景 做法
多项式求值 霍纳法则,\(O(n)\),且不产生大数中间值
求值取模 res = (res * x + c) % mod,每步取模
系数存储 算法用低位在前,输出用高位在前,读入后 reverse 一次定死
多项式乘法 朴素 \(O(nm)\);规模大时用大整数打包(Kronecker)
FFT/NTT 超出笔试范围,本教程不含
判断是否多项式数列 反复差分,\(n\) 次后变常数则为 \(n\) 次多项式
格式化输出 拆成符号/系数/变量三段独立决策,别写巨型 if 链
\(\pm 1\) 省略 只在次数 \(\ge 1\) 时省略,常数项要保留