問題
$p = 998244353$ の有限体上で次数 $N-1$ 以下の $f(x) = \sum a_i x^i$($a_0 = 1$)が与えられる。$g(x)^2 \equiv f(x) \pmod{x^N}$ を満たす $g(x)$($b_0 = 1$)の各係数を出力せよ。
入力形式
N
a_0 a_1 ... a_{N-1}
制約
$1 \le N \le 2^{17} = 131072$
$0 \le a_i < p$
$a_0 = 1$
入出力例
入力例 1
4
1 2 3 4
出力例 1
1 1 1 1
ヒント (段階的開示)
ヒント1: 方向性
$F(g) = g^2 - f = 0$ に FPS Newton 法を適用。$g_{k+1} = \frac{1}{2}(g_k + f/g_k)$ で精度倍々収束。
ヒント2: アプローチ
$g_k^{-1}$ は FPS 逆元(Newton 法)。乗算は NTT で $O(n \log n)$。
ヒント3: 誘導
inv2 = (MOD + 1) // 2
g = [1]
cur = 1
while cur < n:
cur = min(cur * 2, n)
g_inv = fps_inv(g, cur)
fg_inv = ntt_multiply(f[:cur], g_inv)[:cur]
g = [(g[i] + fg_inv[i]) * inv2 % MOD for i in range(cur)]
模範解答 (Python)
import sys
from typing import List
MOD = 998244353
g_primitive_root = 3
def ntt(a: List[int], invert: bool) -> None:
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_primitive_root, (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 = a[i + k]
v = 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 ntt_multiply(a: List[int], b: List[int]) -> List[int]:
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 fps_inv(f: List[int], n: int) -> List[int]:
g = [pow(f[0], MOD - 2, MOD)]
cur = 1
while cur < n:
cur2 = min(cur * 2, n)
fg2 = ntt_multiply(f[:cur2], ntt_multiply(g, g)[:cur2])[:cur2]
two_g = [(2 * gi) % MOD for gi in g] + [0] * (cur2 - len(g))
g = [(two_g[i] - fg2[i]) % MOD for i in range(cur2)]
cur = cur2
return g[:n]
def fps_sqrt(f: List[int], n: int) -> List[int]:
inv2 = (MOD + 1) // 2
g = [1]
cur = 1
while cur < n:
cur = min(cur * 2, n)
g_inv = fps_inv(g, cur)
fg_inv = ntt_multiply(f[:cur], g_inv)[:cur]
new_g = []
for i in range(cur):
gi = g[i] if i < len(g) else 0
fgi = fg_inv[i] if i < len(fg_inv) else 0
new_g.append((gi + fgi) * inv2 % MOD)
g = new_g
return g[:n]
def solve():
data = sys.stdin.read().split()
N = int(data[0])
f = [int(x) for x in data[1:N+1]]
result = fps_sqrt(f, N)
sys.stdout.write(' '.join(map(str, result)) + '\n')
solve()
Step-by-Step 解説
1定式化
$F(g) = g^2 - f = 0$ に Newton 法。$g_{k+1} = g_k - \frac{g_k^2 - f}{2g_k} = \frac{g_k + f g_k^{-1}}{2}$。
$F(g) = g^2 - f = 0$ に Newton 法。$g_{k+1} = g_k - \frac{g_k^2 - f}{2g_k} = \frac{g_k + f g_k^{-1}}{2}$。
2倍精度収束
$g_0 = 1$、精度を $1 \to 2 \to 4 \to \cdots \to N$ と倍々に拡張。$O(\log N)$ ステップ。
$g_0 = 1$、精度を $1 \to 2 \to 4 \to \cdots \to N$ と倍々に拡張。$O(\log N)$ ステップ。
3FPS 逆元
$g_{k+1} = 2g_k - h g_k^2 \pmod{x^{2m}}$。
$g_{k+1} = 2g_k - h g_k^2 \pmod{x^{2m}}$。
4NTT
$p = 998244353 = 119 \cdot 2^{23} + 1$ は NTT 素数。原始根 $g=3$。
$p = 998244353 = 119 \cdot 2^{23} + 1$ は NTT 素数。原始根 $g=3$。
計算量
- fps_sqrt: $O(\log N)$ ステップ × $O(N \log N)$ NTT
- 全体: $O(N \log^2 N)$
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| $f(0) \ne 1$ で $g(0)$ 既定 | 一般には $\sqrt{f(0)} \bmod p$ が必要 | Tonelli-Shanks で二次剰余 |
| NTT サイズ不足 | 結果サイズの 2 冪未満 | 結果サイズの次の 2 冪 |
| fps_inv の初期値 | $f[0] = 0$ で破綻 | $f[0] \ne 0$ を確認 |
inv2 の計算 | 暗算で間違える | (MOD + 1) // 2 |
次のステップ
- $f(0) \ne 1$ の sqrt → Tonelli-Shanks
- FPS の $k$ 乗根 → $\exp(k^{-1} \log f)$
- FPS の log/exp(完全な演算体系)