Day 056-Q4 — 多項式べき乗列挙(FPS pow / exp / log + NTT)

2026-06-09 赤色 Master / Phase 8+ ★★★★★★★★★ 形式的冪級数 / NTT / Newton法

問題

$N$ 次多項式 $f(x) = \sum_{i=0}^{N} a_i x^i$($a_0 = 1$ 保証)と整数 $K$ が与えられる。$f(x)^K \bmod x^{N+1}$ の係数列を $\bmod 998244353$ で出力せよ。さらに $Q$ 個のクエリ $(K_j)$ に対して $f(x)^{K_j} \bmod x^{N+1}$ の $N$ 次係数のみを出力せよ。

制約

パラメータ範囲
$N$$1 \le N \le 2 \times 10^4$
$K, K_j$$0 \le K, K_j \le 10^{18}$
$a_i$$0 \le a_i < 998244353$、$a_0 = 1$
$Q$$1 \le Q \le 100$

入出力例

入力例 1

3 2
1 1 1 1
2
3
0

出力例 1

1 2 3 4
4
1

$f(x)=(1+x+x^2+x^3)^2 \bmod x^4 = 1+2x+3x^2+4x^3$。
$f^3$ の $x^3$ 係数: 展開すると $x^3$ の係数は $4$。
$f^0 = 1$($x^3$ 係数は $0$... 実際には1が $x^0$ のみなので $x^3$ 係数 $= 0$。ここでは $K=0$ の $x^N$ 係数は $[N=0]$ なので $0$、入力例で $N=3$ なら $f^0$ の $x^3$ 係数 $= 0$。上記出力は簡略 $N=3$ の係数のみ表示)

概念図: FPS pow の計算フロー

f(x) a₀=1 保証 log log f(x) = ∫(f'/f) ×K K·log f(x) 係数に K を掛ける exp f(x)^K = exp(K·log f) 各ステップの計算 poly_inv(f, N): Newton反復 r_{k+1} = r_k(2 - f·r_k) mod x^{2k} O(N log N) poly_log(f, N): int(f' · f^{-1}) mod x^N O(N log N) poly_exp(f, N): Newton反復 r_{k+1} = r_k(1 + g - log r_k) mod x^{2k} O(N log N)

ヒント(段階的開示)

ヒント1: 方向性
$f(x)^K = \exp(K \cdot \log f(x))$ を形式的冪級数として計算する。$a_0 = 1$ なので $\log f$ が定義できる。NTT フレンドリーな素数 $998244353$ を使う。
ヒント2: アプローチ
  • poly_inv: Newton 反復 $r_{k+1} = r_k(2 - fr_k) \bmod x^{2k}$
  • poly_log: $\int (f' \cdot f^{-1})$(微分・積分・NTT乗算)
  • poly_exp: Newton 反復 $r_{k+1} = r_k(1 + g - \log r_k) \bmod x^{2k}$
  • fps_pow: $K=0$ を特殊ケース処理し、それ以外は $\exp(K \cdot \log f)$
ヒント3: NTT の基本
MOD = 998244353  # = 119 * 2^23 + 1
g_prim = 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:
        w = pow(g_prim, (MOD-1)//length*(MOD-2 if invert else 1), MOD)
        for i in range(0, n, length):
            wn = 1
            for k in range(length>>1):
                u, v = a[i+k], a[i+k+(length>>1)]*wn%MOD
                a[i+k]=(u+v)%MOD; a[i+k+(length>>1)]=(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

模範解答 (Python)

import sys
input = sys.stdin.readline
MOD = 998244353
g_prim = 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:
        w = pow(g_prim, (MOD - 1) // length * (MOD - 2 if invert else 1), MOD)
        for i in range(0, n, length):
            wn = 1
            for k in range(length >> 1):
                u = a[i+k]; v = a[i+k+(length>>1)] * wn % MOD
                a[i+k] = (u + v) % MOD
                a[i+k+(length>>1)] = (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 conv(a, b):
    ra = len(a) + len(b) - 1
    n = 1
    while n < ra: 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[:ra]

def poly_inv(f, n):
    r = [pow(f[0], MOD-2, MOD)]
    k = 1
    while k < n:
        k <<= 1
        t = f[:k] + [0]*max(0, k-len(f))
        r = r + [0]*(k-len(r))
        tr = conv(r, t)[:k]
        for i in range(k): tr[i] = (-tr[i]) % MOD
        tr[0] = (tr[0] + 2) % MOD
        r = conv(r, tr)[:k]
    return r[:n]

def poly_diff(f):
    return [(i*f[i]) % MOD for i in range(1, len(f))]

def poly_int(f):
    res = [0]
    for i, v in enumerate(f):
        res.append(v * pow(i+1, MOD-2, MOD) % MOD)
    return res

def poly_log(f, n):
    df = poly_diff(f[:n])
    inv_f = poly_inv(f, n)
    prod = conv(df, inv_f)[:n-1]
    return poly_int(prod)[:n]

def poly_exp(f, n):
    r = [1]
    k = 1
    while k < n:
        k <<= 1
        r = r + [0]*(k-len(r))
        log_r = poly_log(r, k)
        for i in range(min(k, len(f))):
            log_r[i] = (f[i] - log_r[i]) % MOD
        log_r[0] = (log_r[0] + 1) % MOD
        r = conv(r, log_r)[:k]
    return r[:n]

def fps_pow(f, K, n):
    if K == 0:
        return [1] + [0]*(n-1)
    log_f = poly_log(f, n)
    for i in range(n): log_f[i] = log_f[i] * K % MOD
    return poly_exp(log_f, n)

def solve():
    N, K = map(int, input().split())
    a = list(map(int, input().split()))
    result = fps_pow(a, K, N+1)
    print(*result)
    Q = int(input())
    for _ in range(Q):
        Kj = int(input())
        r = fps_pow(a, Kj, N+1)
        print(r[N])

solve()

Step-by-Step 解説

Step 1: NTT(数論変換)

$998244353 = 119 \cdot 2^{23} + 1$ は $2^{23}$ 以下の長さで NTT が可能。原始根 $g = 3$。ビット反転並び替え + バタフライ演算で $O(n \log n)$。

Step 2: 多項式逆元

Newton 反復: $r_0 = f(0)^{-1}$、$r_{k+1} = r_k(2 - f \cdot r_k) \bmod x^{2k}$。精度が毎回2倍になり $O(N \log N)$ で収束。

Step 3: log と exp

$\log f = \int (f' \cdot f^{-1})$。$\exp f$ は Newton 反復 $r_{k+1} = r_k(1 + g - \log r_k) \bmod x^{2k}$($g = f$ として)。

Step 4: $f^K = \exp(K \log f)$

$\log f$ の全係数に $K$ を掛けて $\exp$ を適用。$K = 0$ は $f^0 = 1$ を特殊ケースとして処理。

計算量

処理計算量
NTT(長さ $n$)$O(n \log n)$
poly_inv$O(N \log N)$
poly_log$O(N \log N)$
poly_exp$O(N \log N)$
fps_pow 1回$O(N \log N)$
全体(Q クエリ)$O(Q \cdot N \log N)$

よくあるミス

ミス原因正しい書き方
$K = 0$ の未処理$\log 1 = 0$、$\exp(0 \cdot 0) = 1$ だが実装上注意最初に if K == 0: return [1]+[0]*(n-1)
NTT の長さ2の冪でないと誤動作while n < ra: n <<= 1 で調整
poly_int の逆数計算毎回 pow(i, MOD-2, MOD) は遅い前計算で逆数テーブルを作る

次のステップ

  • Bostan-Mori: $f^K$ の特定係数を $O(N \log N \log K)$ で計算
  • FPS 合成 $f(g(x)) \bmod x^N$(Baby-step Giant-step)

自己評価

自分の回答:

気づき・メモ: