Day 091-Q4 — NTT(数論変換・多項式乗算 $O(N\log N)$)

2026-07-14 赤色 Master / Phase 8+ ★★★★★★★★★ NTT・畳み込み・998244353

問題

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$。

概念図

係数表現 → 点値表現で掛け算 → 係数表現へ戻す A(x), B(x)係数 点値 Â, B̂NTT 順変換 Ĉ = ·B̂点ごとの積 C(x)逆変換 998244353 = 119·2^23 + 1 は NTT 素数(原始根 3)

ヒント

ヒント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)

自己評価

理解度: / /

自分の回答:

気づき・メモ: