Day 080-Q4 — Strassen 法による行列乗算($O(N^{\log_2 7}) \approx O(N^{2.807})$)

2026-07-03 赤色 Master / Phase 8+ ★★★★★★★★★ 数値線形代数・Strassen・再帰分割・高速行列積

問題

$N \times N$ 行列 $A$ と $B$ の積 $C = AB$ を、乗算回数 $O(N^{\log_2 7}) \approx O(N^{2.807})$ のアルゴリズムで計算せよ。

全要素は $\text{mod } (10^9 + 7)$ で与えられる。

出力: 行列 $C = AB \pmod{10^9 + 7}$ の全要素を左上から行優先で出力。

制約

パラメータ範囲備考
$N$$N = 2^k$($1 \le k \le 8$)つまり $N \le 256$
行列要素$0 \le a_{ij}, b_{ij} < 10^9 + 7$mod p

入出力例

入力例1(2×2)

2
1 2
3 4
5 6
7 8

出力例1

19 22
43 50

概念図: Strassen の7積分解

Strassen法: 2×2ブロック行列を7回の乗算で計算 入力行列の分割 A₁₁ A₁₂ A₂₁ A₂₂ × B₁₁ B₁₂ B₂₁ B₂₂ = C₁₁ C₁₂ C₂₁ C₂₂ 7つの積 M₁…M₇ M₁ = (A₁₁+A₂₂)(B₁₁+B₂₂) M₂ = (A₂₁+A₂₂)B₁₁ M₃ = A₁₁(B₁₂-B₂₂) M₄ = A₂₂(B₂₁-B₁₁) M₅ = (A₁₁+A₁₂)B₂₂ M₆ = (A₂₁-A₁₁)(B₁₁+B₁₂) M₇ = (A₁₂-A₂₂)(B₂₁+B₂₂) C ブロックの組み立て(加減算のみ) C₁₁ = M₁ + M₄ - M₅ + M₇ C₁₂ = M₃ + M₅ C₂₁ = M₂ + M₄ C₂₂ = M₁ - M₂ + M₃ + M₆ 計算量比較 通常の行列積: T(N) = 8·T(N/2) + O(N²) → O(N³) Strassen: T(N) = 7·T(N/2) + O(N²) → O(N^{log₂7}) ≈ O(N^{2.807}) N=256 のとき: 通常 ≈ 1.68×10⁷ vs Strassen ≈ 8.3×10⁶ 再帰: N→N/2→N/4→…→THRESHOLD(32以下で通常積)

ヒント

ヒント1(方向性)

通常の行列積は $O(N^3)$。Strassen(1969)は $2\times2$ ブロック行列の積を 7回の再帰乗算 で実現し、$O(N^{2.807})$ を達成した。

ヒント2(アプローチ)

$2N \times 2N$ 行列を 4 つの $N \times N$ ブロックに分割し、7 つの積 $M_1, \ldots, M_7$ を定義して、$C$ の各ブロックを加減算のみで組み立てる。

$T(N) = 7 T(N/2) + O(N^2)$ → マスター定理より $T(N) = O(N^{\log_2 7})$

ヒント3(ほぼ答え)
M1 = strassen(mat_add(A11, A22), mat_add(B11, B22))
M2 = strassen(mat_add(A21, A22), B11)
M3 = strassen(A11,               mat_sub(B12, B22))
M4 = strassen(A22,               mat_sub(B21, B11))
M5 = strassen(mat_add(A11, A12), B22)
M6 = strassen(mat_sub(A21, A11), mat_add(B11, B12))
M7 = strassen(mat_sub(A12, A22), mat_add(B21, B22))

C11 = mat_add(mat_sub(mat_add(M1, M4), M5), M7)
C12 = mat_add(M3, M5)
C21 = mat_add(M2, M4)
C22 = mat_add(mat_add(mat_sub(M1, M2), M3), M6)

模範解答

import sys
input = sys.stdin.readline
MOD = 10**9 + 7

def mat_add(A, B):
    n = len(A)
    return [[(A[i][j] + B[i][j]) % MOD for j in range(n)] for i in range(n)]

def mat_sub(A, B):
    n = len(A)
    return [[(A[i][j] - B[i][j]) % MOD for j in range(n)] for i in range(n)]

def mat_mul_naive(A, B):
    n = len(A)
    C = [[0]*n for _ in range(n)]
    for i in range(n):
        for k in range(n):
            if A[i][k] == 0: continue
            for j in range(n):
                C[i][j] = (C[i][j] + A[i][k] * B[k][j]) % MOD
    return C

def split(M):
    n, h = len(M), len(M) // 2
    A11 = [row[:h] for row in M[:h]]
    A12 = [row[h:] for row in M[:h]]
    A21 = [row[:h] for row in M[h:]]
    A22 = [row[h:] for row in M[h:]]
    return A11, A12, A21, A22

def join(C11, C12, C21, C22):
    n = len(C11)
    C = [C11[i] + C12[i] for i in range(n)]
    C += [C21[i] + C22[i] for i in range(n)]
    return C

THRESHOLD = 32

def strassen(A, B):
    n = len(A)
    if n <= THRESHOLD:
        return mat_mul_naive(A, B)
    A11, A12, A21, A22 = split(A)
    B11, B12, B21, B22 = split(B)

    M1 = strassen(mat_add(A11, A22), mat_add(B11, B22))
    M2 = strassen(mat_add(A21, A22), B11)
    M3 = strassen(A11,               mat_sub(B12, B22))
    M4 = strassen(A22,               mat_sub(B21, B11))
    M5 = strassen(mat_add(A11, A12), B22)
    M6 = strassen(mat_sub(A21, A11), mat_add(B11, B12))
    M7 = strassen(mat_sub(A12, A22), mat_add(B21, B22))

    C11 = mat_add(mat_sub(mat_add(M1, M4), M5), M7)
    C12 = mat_add(M3, M5)
    C21 = mat_add(M2, M4)
    C22 = mat_add(mat_add(mat_sub(M1, M2), M3), M6)

    return join(C11, C12, C21, C22)

def main():
    N = int(input())
    A = [list(map(int, input().split())) for _ in range(N)]
    B = [list(map(int, input().split())) for _ in range(N)]
    C = strassen(A, B)
    for row in C:
        print(*row)

main()

Step-by-Step 解説

Step 1: 通常の行列積の限界

$N \times N$ 行列積は $O(N^3)$ で乗算 $N^3$ 回。$N = 1000$ で $10^9$ 回は TLE。

Step 2: Strassen の発見

$2 \times 2$ ブロックの積を 8 回から 7 回の乗算 に削減。再帰関係 $T(N) = 7 T(N/2) + O(N^2)$ → マスター定理より $T(N) = O(N^{\log_2 7})$。

Step 3: 7つの積の定式化

$$M_1 = (A_{11}+A_{22})(B_{11}+B_{22}), \quad M_2 = (A_{21}+A_{22})B_{11}$$ $$M_3 = A_{11}(B_{12}-B_{22}), \quad M_4 = A_{22}(B_{21}-B_{11})$$ $$M_5 = (A_{11}+A_{12})B_{22}, \quad M_6 = (A_{21}-A_{11})(B_{11}+B_{12})$$ $$M_7 = (A_{12}-A_{22})(B_{21}+B_{22})$$

Step 4: 結果の組み立て

$$C_{11} = M_1 + M_4 - M_5 + M_7, \quad C_{12} = M_3 + M_5$$ $$C_{21} = M_2 + M_4, \quad C_{22} = M_1 - M_2 + M_3 + M_6$$

Step 5: 閾値の設定

再帰の葉では THRESHOLD 以下で通常乗算に切り替える。Python では $32 \sim 64$ が適切(再帰オーバーヘッドの考慮)。

よくあるミス

ミス原因正しい書き方
乗算回数が8回になる$M_i$ の定義ミス教科書の式を正確に写す
join での行の順序ブロックを結合するとき上下を逆にするC[:h] = C11, C[h:] = C21
Threshold を1にするPython では再帰オーバーヘッドが大THRESHOLD = 32 以上に設定
負の値が生じるmod 演算で引き算が負% MOD を必ず適用

次のステップ

  • 発展問題: Winograd 変換(乗算 6 回・加算増加)
  • 現状最速: Coppersmith-Winograd $O(N^{2.372})$、Williams $O(N^{2.371...})$

自己評価

理解度: / /

自分の回答:

気づき・メモ: