問題
2つの多項式 $A(x)=\sum a_i x^i$(次数 $<N$)と $B(x)=\sum b_j x^j$(次数 $<M$)の積 $C(x)=A(x)B(x)$ の全係数を $998244353$ で割った余りで出力する。
制約
| パラメータ | 範囲 | 備考 |
|---|---|---|
| $N, M$ | $1 \le N, M \le 2\times10^5$ | 各多項式の項数 |
| $a_i, b_j$ | $0 \le a_i, b_j < 998244353$ | 係数 |
| 法 | $998244353$ | NTT 素数・原始根 3 |
入出力例
入力例1
2 2
1 2
3 4
出力例1
3 10 8
$(1+2x)(3+4x)=3+10x+8x^2$。
入力例2
3 1
1 2 3
5
出力例2
5 10 15
$(1+2x+3x^2)\cdot5=5+10x+15x^2$。
概念図
ヒント
ヒント1(方向性)
素朴な畳み込みは $O(NM)$。法 $998244353$ は NTT 素数なので FFT の整数版 NTT が使える。
ヒント2(アプローチ)
長さ $2$ 冪 $n\ge N+M-1$ に0埋め、順変換で点値表現に→各点で掛ける→逆変換で係数表現に戻す。
ヒント3(ほぼ答え)
w = pow(3, (MOD-1)//length, MOD) # 順変換
w = pow(3, MOD-1-(MOD-1)//length, MOD) # 逆変換(w の逆元)
# 逆変換の最後に全要素へ n^{-1} を掛ける
模範解答
import sys
MOD = 998244353
ROOT = 3
def ntt(a, invert):
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:
if invert:
w = pow(ROOT, MOD - 1 - (MOD - 1) // length, MOD)
else:
w = pow(ROOT, (MOD - 1) // length, MOD)
half = length >> 1
for i in range(0, n, length):
wn = 1
for k in range(half):
u = a[i + k]
v = a[i + k + half] * wn % MOD
a[i + k] = (u + v) % MOD
a[i + k + half] = (u - v) % MOD
wn = wn * w % MOD
length <<= 1
if invert:
inv_n = pow(n, MOD - 2, MOD)
for i in range(n):
a[i] = a[i] * inv_n % MOD
def multiply(a, b):
result_size = len(a) + len(b) - 1
n = 1
while n < result_size:
n <<= 1
fa = a + [0] * (n - len(a))
fb = b + [0] * (n - len(b))
ntt(fa, False)
ntt(fb, False)
for i in range(n):
fa[i] = fa[i] * fb[i] % MOD
ntt(fa, True)
return fa[:result_size]
def main():
data = sys.stdin.buffer.read().split()
idx = 0
n = int(data[idx]); m = int(data[idx + 1]); idx += 2
a = [int(data[idx + i]) % MOD for i in range(n)]; idx += n
b = [int(data[idx + i]) % MOD for i in range(m)]; idx += m
c = multiply(a, b)
sys.stdout.write(' '.join(map(str, c)) + '\n')
main()
Step-by-Step 解説
Step 1: 長さを2の冪に揃える
結果係数は $N+M-1$ 個。これ以上の最小の2の冪 $n$ に0埋めする。
Step 2: ビット反転並べ替え
in-place バタフライのため入力をビット反転順に並べ替える($O(n)$)。
Step 3: バタフライ演算(順変換)
段長 length を $2,4,8,\dots$ と倍にしつつ、回転因子 $w=g^{(p-1)/\text{length}}$($g=3$)で合成する。
Step 4: 点ごとの積と逆変換
点値表現を各点で掛け、逆変換(回転因子を逆元にし末尾で $n^{-1}$ 倍)で係数へ戻す。先頭 $N+M-1$ 個が答え。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| 長さを $N+M-1$ のまま使う | NTT は2冪長が前提 | 最小の2の冪に0埋め |
| 逆変換で $n^{-1}$ 掛け忘れ | 正規化漏れ | 末尾で全要素に $n^{-1}$ |
| 任意素数で原始根を使う | NTT 素数が必要 | $998244353$(原始根3)等 |
| 減算で負のまま出力 | % の符号 | (u-v)%MOD は非負 |
次のステップ
- 発展: 任意 mod での畳み込み(3素数 NTT + Garner/CRT)
- 発展: FPS の逆元・log・exp を NTT + Newton 法で
- 次回予告: 2-SAT(含意グラフ + SCC)
自己評価
理解度: / /
自分の回答:
気づき・メモ: