畳み込みを「係数の組」で理解する

配列 A,B の畳み込み C[k] は i+j=k のすべての A[i]B[j] の和です。数え上げと多項式の積を結び付けます。

言語:Python / 計算量:基準 O(NM)、NTT O(L log L)

前提:剰余と組合せを安全に扱う / 行列累乗で線形遷移を飛ばす

考え方

  1. まず O(NM) の二重ループで定義通りに実装。
  2. 出力長は N+M-1。
  3. NTT は適した法と原始根のもとで O(L log L) へ高速化します。

具体例

[1,2] と [3,4] の積は [3,10,8]。中央は1×4+2×3です。

実装

a=[1,2];b=[3,4]
c=[0]*(len(a)+len(b)-1)
for i,x in enumerate(a):
    for j,y in enumerate(b):c[i+j]+=x*y
print(c)

注意する条件

法と変換長に条件があります。998244353ではL≤2^23。Pythonでの実用上限は時間・メモリ次第です。講義後半にNTT実装。

確認問題

[1,2] と [3,4] の畳み込みの中央係数は?

解答と理由

10

1×4 + 2×3 = 10です。

実装課題

N M、配列 A、配列 B が3行。畳み込みを998244353で割った余りで出す。1≤N,M≤100。まず基準実装を作る。

入力:
2 2
1 2
3 4
出力:
3 10 8
参考実装
n,m=map(int,input().split());a=list(map(int,input().split()));b=list(map(int,input().split()))
p=998244353;c=[0]*(n+m-1)
for i,x in enumerate(a):
    for j,y in enumerate(b):c[i+j]=(c[i+j]+x*y)%p
print(*c)

998244353でのNTT実装

長さLを必要な出力長以上の2の冪にし、配列を0で埋めます。998244353 = 119×2^23+1 なので、ここでは L≤2^23 が必要です。Pythonでこの上限まで実用的に扱えるという意味ではなく、メモリと実行時間は別に見積もります。

順変換で多項式をL個の点での値へ変え、点ごとに掛け、逆変換で係数へ戻します。段階ごとに偶数側と奇数側を組み合わせるバタフライ演算を行うことで O(L log L) にします。

逆変換では根を逆元へ変え、最後にLの逆元を全体へ掛けます。定数倍は小さくないため、大きい制約ではC++のACLも比較してください。まず小さい乱数列を二重ループの答えと照合しましょう。

MOD = 998244353
ROOT = 3

def ntt(a, invert=False):
    n = len(a)
    assert n > 0 and n & (n-1) == 0 and n <= (1 << 23)
    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:
        root = pow(ROOT, (MOD-1)//length, MOD)
        if invert:
            root = pow(root, MOD-2, MOD)
        half = length//2
        for start in range(0, n, length):
            w = 1
            for i in range(start, start+half):
                u = a[i]
                v = a[i+half]*w % MOD
                a[i] = (u+v) % MOD
                a[i+half] = (u-v) % MOD
                w = w*root % MOD
        length *= 2
    if invert:
        inv = pow(n, MOD-2, MOD)
        for i in range(n):
            a[i] = a[i]*inv % MOD

def convolution(a, b):
    if not a or not b:
        return []
    size = 1
    while size < len(a)+len(b)-1:
        size *= 2
    assert size <= (1 << 23)
    x = [v % MOD for v in a] + [0]*(size-len(a))
    y = [v % MOD for v in b] + [0]*(size-len(b))
    ntt(x)
    ntt(y)
    for i in range(size):
        x[i] = x[i]*y[i] % MOD
    ntt(x, True)
    return x[:len(a)+len(b)-1]

print(convolution([1,2], [3,4]))  # [3, 10, 8]

読了の記録・下書き・メモへ

関連する公式資料・課題