問題
次数 $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。
Horner で $O(NM)$、$2 \times 10^5$ では TLE。
2Remainder Tree
$f(x_i)$ = $f \bmod (x-x_i)$ の定数項。分割統治で剰余を降ろす。
$f(x_i)$ = $f \bmod (x-x_i)$ の定数項。分割統治で剰余を降ろす。
3計算量解析
Tree 構築 / 剰余降ろし ともに $O(M \log^2 M)$。
Tree 構築 / 剰余降ろし ともに $O(M \log^2 M)$。
4NTT 高速化
多項式剰余を Newton 法で $O(n \log n)$ にすると更に高速。
多項式剰余を Newton 法で $O(n \log n)$ にすると更に高速。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| size を M にする | NTT が壊れる | 2 の冪に切り上げ |
| 空多項式 | 単位元忘れ | 余りスロットに [1] |
| poly_mod 終了条件 | 先頭 0 のまま比較 | 後ろの 0 を strip |
次のステップ
- 多項式補間(Polynomial Interpolation)