Day 028-Q4 — 部分集合畳み込み(Walsh-Hadamard変換 + Subset Sum Convolution)

2026-05-11 赤色 Master / Phase 8+ ★★★★★★★★★ XOR畳み込み応用

問題

長さ $2^N$ の整数列 $f(0), f(1), \ldots, f(2^N - 1)$ が与えられる。以下の畳み込みを計算せよ。

$$g(S) = \sum_{T \subseteq S} f(T) \cdot f(S \setminus T)$$

ただし $S$ はビットマスク($0$ から $2^N - 1$)、$T \subseteq S$ は $S$ のすべての部分集合。この計算を $O(3^N)$ より速く、$O(N^2 \cdot 2^N)$ で行え。

入力形式

N
f_0 f_1 ... f_{2^N - 1}

制約

$1 \le N \le 20$
$0 \le f_i \le 10^9$
答えは $10^{18}$ 以下

入出力例

入力例 1

3
1 2 3 4 5 6 7 8

出力例 1

2 8 10 56 14 76 84 512

出力値は実装で確認すること。

ヒント (段階的開示)

ヒント1: 方向性
部分集合畳み込みは「要素数(popcount)を固定した上でゼータ変換・逆変換」を行う。$f$ をランクごとに分離。
ヒント2: アプローチ
$\hat{f}[k][S] = \sum_{T \subseteq S, |T|=k} f(T)$ と定義。$\hat{g}[k][S] = \sum_{j=0}^{k} \hat{f}[j][S] \cdot \hat{f}[k-j][S]$ で計算し、逆ゼータ変換で $g$ を得る。
ヒント3: 誘導
ranked = [[0] * (1 << N) for _ in range(N + 1)]
for s in range(1 << N):
    ranked[popcount(s)][s] = f[s]
# SOS変換 / ランク畳み込み / 逆メビウス変換

模範解答 (Python)

import sys
input = sys.stdin.readline

def solve():
    N = int(input())
    f = list(map(int, input().split()))

    size = 1 << N

    def popcount(x):
        return bin(x).count('1')

    # Step 1: ランク付き配列の初期化
    ranked = [[0] * size for _ in range(N + 1)]
    for s in range(size):
        ranked[popcount(s)][s] = f[s]

    # Step 2: 各ランクについて SOS ゼータ変換
    for k in range(N + 1):
        for i in range(N):
            for s in range(size):
                if (s >> i) & 1:
                    ranked[k][s] += ranked[k][s ^ (1 << i)]

    # Step 3: ランクの畳み込み
    conv = [[0] * size for _ in range(N + 1)]
    for k in range(N + 1):
        for j in range(k + 1):
            for s in range(size):
                conv[k][s] += ranked[j][s] * ranked[k - j][s]

    # Step 4: 逆 SOS 変換(メビウス変換)
    for k in range(N + 1):
        for i in range(N):
            for s in range(size):
                if (s >> i) & 1:
                    conv[k][s] -= conv[k][s ^ (1 << i)]

    # Step 5: g[s] = conv[popcount(s)][s]
    g = [0] * size
    for s in range(size):
        g[s] = conv[popcount(s)][s]

    print(*g)

solve()

Step-by-Step 解説

1問題の難しさ
素朴に計算すると $O(3^N)$。$N=20$ で $3^{20} \approx 3.5 \times 10^9$ で TLE。
2ランク分解
$T$ と $S\setminus T$ が disjoint な組を集計したいが、直接ゼータ変換すると不素な組も混入。$f$ をビット数(ランク)ごとに分離して回避。
3ランク付きゼータ変換
各ランク $k$ について SOS DP を適用。$N+1$ 個のランク × $O(N \cdot 2^N)$。
4点ごとのランク畳み込み
各 $S$ で多項式積を計算。$\hat{g}^{(k)}(S) = \sum_j \hat{f}^{(j)}(S) \cdot \hat{f}^{(k-j)}(S)$。
5逆変換と取り出し
逆メビウス変換で $g^{(k)}$ を復元し、$g(S) = g^{(|S|)}(S)$。

計算量

  • ランク付きゼータ変換: $O(N^2 \cdot 2^N)$
  • ランク畳み込み: $O(N^2 \cdot 2^N)$
  • 逆ゼータ変換: $O(N^2 \cdot 2^N)$
  • 全体: $O(N^2 \cdot 2^N)$($N=20$ で約 $4.2 \times 10^8$、PyPy 推奨)

よくあるミス

ミス原因正しい書き方
通常の SOS DP との混同ランクを区別しないranked[k][s] と2次元で管理
ゼータ変換の方向上位 vs 下位が不明確部分集合和は「0→1の方向」で加算
逆変換の符号引き算の方向を誤るconv[k][s] -= conv[k][s^(1<<i)]
g[s] の取り出しk を間違えるk = popcount(s)
計算量の見誤り$O(N^2 2^N)$ になっていない3重ループ構造を確認

次のステップ

  • 集合被覆カウント(彩色多項式)
  • OR/AND 畳み込みとの比較
  • $g(S) = \sum_{T \cap S = \emptyset} f(T)$ との違い

自己評価

自分の回答

気づき・メモ