Day 059-Q1 — XOR畳み込み (Walsh-Hadamard変換)

2026-06-12 赤色 Master / Phase 8+ ★★★★★★★★★ XOR畳み込み / WHT / 部分集合XOR二乗和

問題

$N$ 個の非負整数 $a_0, a_1, \dots, a_{N-1}$($0 \le a_i < 2^M$)が与えられる。次の値を $\mathrm{mod}\ 998244353$ で求めよ。

$$\sum_{S \subseteq \{0,\dots,N-1\}} \left(\bigoplus_{i \in S} a_i\right)^2$$

ただし $\bigoplus$ はビット XOR、空集合の XOR は $0$ とする。

制約

パラメータ範囲
$N$$1 \le N \le 2 \times 10^5$
$M$$1 \le M \le 20$
$a_i$$0 \le a_i < 2^M$

入出力例

入力例 1

3 2
1 2 3

出力例 1

30

$a = [1,2,3]$。部分集合 $\emptyset(0), \{1\}(1), \{2\}(2), \{3\}(3), \{1,2\}(3), \{1,3\}(2), \{2,3\}(1), \{1,2,3\}(0)$ の XOR の2乗和 = $0+1+4+9+9+4+1+2 = 30$。

概念図: Walsh-Hadamard 変換

WHT の蝶演算: 各ステップで h 幅の対を (x+y, x-y) に変換 配列 a = [a0, a1, a2, a3] (M=2, SIZE=4) a0 a1 a2 a3 h=1: ペア (0,1), (2,3) を蝶演算 x+y, x-y a0+a1 a0-a1 a2+a3 a2-a3 h=2: ペア (0,2), (1,3) を蝶演算 → 最終 WHT â0 â1 â2 â3 XOR 畳み込みの流れ 1. f = [0]*SIZE; f[0] = 1 (空集合) 2. WHT(f) 3. for each a_i: delta = [0]*SIZE; delta[a_i] = 1 WHT(delta) f[i] *= (delta[i] + 1) ← 点ごと積 4. IWHT(f) → f[v] = 部分集合XOR=vの個数 5. ans = Σ v² * f[v] 計算量: O(N·2^M + 2^M·M) N=2×10⁵, M=20 → 約 4×10⁷ 回

ヒント(段階的開示)

ヒント1: 方向性
全部分集合の XOR の2乗和を直接計算すると $O(2^N)$。これを XOR 畳み込み(Walsh-Hadamard 変換)と部分集合の統計で $O(N \cdot 2^M + 2^M \cdot M)$ に落とせる。まず「各値 $v$ に対して全部分集合 XOR が $v$ になる個数 $f[v]$」を効率的に求めることを考えよ。
ヒント2: アプローチ
  • $f[v] = |\{S \subseteq [N] : \bigoplus_{i \in S} a_i = v\}|$ を求める
  • 空集合から始め、各要素 $a_i$ を加えるたびに f'[v XOR a_i] += f[v]
  • これは $f$ と $\mathbf{1}_{a_i}$($a_i$ 番だけ1のベクトル)の XOR 畳み込み
  • WHT ドメインで「$(1 + \hat{\delta}_{a_i})$」を全 $N$ 要素分掛け算
ヒント3: コード骨格
def wht(a, inv=False):
    n = len(a)
    h = 1
    while h < n:
        for i in range(0, n, h*2):
            for j in range(h):
                x, y = a[i+j], a[i+j+h]
                a[i+j], a[i+j+h] = (x+y)%MOD, (x-y)%MOD
        h <<= 1
    if inv:
        inv_n = pow(n, MOD-2, MOD)
        for i in range(n): a[i] = a[i]*inv_n%MOD

f = [0]*SIZE; f[0] = 1
wht(f)
for a in A:
    delta = [0]*SIZE; delta[a] = 1
    wht(delta)
    for i in range(SIZE): f[i] = f[i]*(delta[i]+1)%MOD
wht(f, inv=True)
ans = sum(v*v*f[v] for v in range(SIZE)) % MOD

模範解答 (Python)

import sys
input = sys.stdin.readline

MOD = 998244353

def wht(a, inv=False):
    n = len(a)
    h = 1
    while h < n:
        for i in range(0, n, h*2):
            for j in range(h):
                x, y = a[i+j], a[i+j+h]
                a[i+j], a[i+j+h] = (x+y)%MOD, (x-y)%MOD
        h <<= 1
    if inv:
        inv_n = pow(n, MOD-2, MOD)
        for i in range(n):
            a[i] = a[i]*inv_n%MOD

def main():
    N, M = map(int, input().split())
    A = list(map(int, input().split()))
    SIZE = 1 << M

    f = [0]*SIZE
    f[0] = 1  # 空集合
    wht(f)

    for a in A:
        delta = [0]*SIZE
        delta[a] = 1
        wht(delta)
        for i in range(SIZE):
            f[i] = f[i]*(delta[i]+1)%MOD

    wht(f, inv=True)

    ans = 0
    for v in range(SIZE):
        ans = (ans + v*v%MOD*f[v]) % MOD
    print(ans)

main()

Step-by-Step 解説

Step 1: 問題の変換

$$\sum_S (\mathrm{XOR}_S)^2 = \sum_{v=0}^{2^M-1} v^2 \cdot f[v]$$ ここで $f[v] = |\{S : \mathrm{XOR}_S = v\}|$。$f$ を求めることが目標。

Step 2: XOR 畳み込みによる $f$ の更新

初期 $f = \mathbf{1}_0$(空集合のみ)。各要素 $a_i$ を加えると: $f^{\text{new}}[v] = f[v] + f[v \oplus a_i]$($a_i$ を選ぶ/選ばない)。 これを全要素に適用 = $f$ と $\mathbf{1}_{a_i}$ の XOR 畳み込みの "加算版"。

Step 3: WHT の活用

WHT ドメインでは XOR 畳み込みが点ごとの積に変換される。$\hat{f} \cdot (1 + \hat{\delta}_{a_i})$ を全要素に適用後、逆 WHT で $f$ を復元。

Step 4: 計算量

WHT: $O(2^M \cdot M)$。全要素の delta 適用: $O(N \cdot 2^M)$。合計 $O(N \cdot 2^M + 2^M \cdot M)$。$N = 2 \times 10^5, M = 20$ で約 $4 \times 10^7$ 演算。

よくあるミス

ミス原因正しい書き方
逆 WHT で inv_n = pow(2, MOD-2, MOD)$n = 2^M$ 全体の逆数が必要inv_n = pow(n, MOD-2, MOD)
空集合を忘れるf[0]=0 のままf[0] = 1
f[i] *= delta[i] のみ「選ばない」場合(元の $f$)を消すf[i] *= (delta[i] + 1)

次のステップ

発展問題: $\sum_S (\mathrm{XOR}_S = T)$ を満たす部分集合数を全 $T$ について求めよ(XOR Basis 応用との比較)。

自己評価

解いた後に記入してください。