Day 031-Q2 — 多項式剰余・CRT補間(Multipoint Evaluation + Remainder Tree)

2026-05-14 赤色 / Phase 8+ ★★★★★★★★★ NTT + 多項式剰余木

問題

次数 $N-1$ の多項式 $f$ と $M$ 点 $x_i$ に対し、$f(x_i) \bmod p$ を全部求めよ。

制約

$1 \le N, M \le 2 \times 10^5$
$p = 998244353$
$0 \le a_i, x_i < p$

入出力例

入力例 1

3 4
1 2 3
0 1 2 3

出力例 1

1
6
17
34

ヒント (段階的開示)

ヒント1: 方向性
単純 Horner は $O(NM)$ で TLE。Remainder Tree で $O((N+M) \log^2(N+M))$。
ヒント2: アプローチ
$P(x) = \prod (x - x_i)$ を分割統治で構築 → $f \bmod P$ を再帰的に降ろす。葉で $f(x_i) = f \bmod (x - x_i)$。
ヒント3: NTT
$p = 998244353$ は NTT 素数。原始根 $g=3$。

模範解答 (Python)

import sys
from sys import stdin

MOD = 998244353
g = 3

def ntt(a, invert=False):
    n = len(a)
    j = 0
    for i in range(1, n):
        bit = n >> 1
        while j & bit:
            j ^= bit
            bit >>= 1
        j ^= bit
        if i < j:
            a[i], a[j] = a[j], a[i]
    length = 2
    while length <= n:
        w = pow(g, (MOD - 1) // length, MOD)
        if invert:
            w = pow(w, MOD - 2, MOD)
        for i in range(0, n, length):
            wn = 1
            for k in range(length // 2):
                u, v = a[i+k], a[i+k+length//2] * wn % MOD
                a[i+k] = (u + v) % MOD
                a[i+k+length//2] = (u - v) % MOD
                wn = wn * w % MOD
        length <<= 1
    if invert:
        n_inv = pow(n, MOD - 2, MOD)
        for i in range(n):
            a[i] = a[i] * n_inv % MOD

def poly_mul(a, b):
    if not a or not b:
        return []
    result_len = len(a) + len(b) - 1
    n = 1
    while n < result_len:
        n <<= 1
    fa = a + [0] * (n - len(a))
    fb = b + [0] * (n - len(b))
    ntt(fa); ntt(fb)
    fc = [fa[i] * fb[i] % MOD for i in range(n)]
    ntt(fc, invert=True)
    return fc[:result_len]

def poly_mod(a, b):
    a = list(a); b = list(b)
    while len(a) >= len(b):
        if a[-1] == 0:
            a.pop(); continue
        coef = a[-1] * pow(b[-1], MOD - 2, MOD) % MOD
        deg_diff = len(a) - len(b)
        for i in range(len(b)):
            a[deg_diff + i] = (a[deg_diff + i] - coef * b[i]) % MOD
        while a and a[-1] == 0:
            a.pop()
    return a

def multipoint_eval(f, xs):
    M = len(xs)
    if M == 0:
        return []
    size = 1
    while size < M:
        size <<= 1
    tree = [None] * (2 * size)
    for i in range(M):
        tree[size + i] = [(-xs[i]) % MOD, 1]
    for i in range(M, size):
        tree[size + i] = [1]
    for i in range(size - 1, 0, -1):
        tree[i] = poly_mul(tree[2*i], tree[2*i+1])
    remainder = [None] * (2 * size)
    remainder[1] = poly_mod(f, tree[1])
    for i in range(1, size):
        if remainder[i] is None:
            continue
        if tree[2*i]:
            remainder[2*i] = poly_mod(remainder[i], tree[2*i])
        if tree[2*i+1]:
            remainder[2*i+1] = poly_mod(remainder[i], tree[2*i+1])
    result = []
    for i in range(M):
        rem = remainder[size + i]
        result.append(rem[0] % MOD if rem else 0)
    return result

def solve():
    data = stdin.read().split()
    idx = 0
    N, M = int(data[idx]), int(data[idx+1]); idx += 2
    a = [int(data[idx+i]) % MOD for i in range(N)]; idx += N
    xs = [int(data[idx+i]) % MOD for i in range(M)]; idx += M
    results = multipoint_eval(a, xs)
    print('\n'.join(map(str, results)))

solve()

Step-by-Step 解説

1ナイーブの計算量
Horner で $O(NM)$、$2 \times 10^5$ では TLE。
2Remainder Tree
$f(x_i)$ = $f \bmod (x-x_i)$ の定数項。分割統治で剰余を降ろす。
3計算量解析
Tree 構築 / 剰余降ろし ともに $O(M \log^2 M)$。
4NTT 高速化
多項式剰余を Newton 法で $O(n \log n)$ にすると更に高速。

よくあるミス

ミス原因正しい書き方
size を M にするNTT が壊れる2 の冪に切り上げ
空多項式単位元忘れ余りスロットに [1]
poly_mod 終了条件先頭 0 のまま比較後ろの 0 を strip

次のステップ

  • 多項式補間(Polynomial Interpolation)

自己評価

自分の回答

気づき・メモ