Matrix-chain Multiplication

My vault 演算法筆記:Matrix-chain Multiplication。

1. 問題定義 (Problem)

給定一串矩陣鏈 $\langle A_1, A_2, \dots, A_n \rangle$,找出一個最佳的「加括號方式」(parenthesization),使得計算總乘積 $A_1A_2\dots A_n$ 所需的「純量乘法次數」最少。

  • 關鍵: 矩陣乘法有結合律,但沒有交換律。

  • 成本: $A_{p \times q} \times B_{q \times r}$ 的成本是 $p \times q \times r$ 次純量乘法。

為什麼這很重要?

不同的計算順序,成本天差地遠。

範例: $A_1 (10 \times 100)$, $A_2 (100 \times 5)$, $A_3 (5 \times 50)$

  • 順序 1: $((A_1A_2)A_3)$

    • $(A_1A_2)$: $10 \times 100 \times 5 = 5,000$

    • $(…)A_3$: $10 \times 5 \times 50 = 2,500$

    • 總成本: 7,500

  • 順序 2: $(A_1(A_2A_3))$

    • $(A_2A_3)$: $100 \times 5 \times 50 = 25,000$

    • $A_1(…)$: $10 \times 100 \times 50 = 50,000$

    • 總成本: 75,000

2. 為什麼不能用暴力法 (Brute Force)?

暴力法需要嘗試所有可能的加括號方式。

  • $P(n)$:$n$ 個矩陣的加括號方式數量。

  • $P(n) = \sum_{k=1}^{n-1} P(k)P(n-k)$

  • 這個數列是卡特蘭數 (Catalan numbers),呈指數級增長 ($\Omega(4n / n{3/2})$)。

  • 當 $n$ 很大時,暴力法不可行。

3. 動態規劃 (DP) 解法

這是一個 Interval DP 的經典問題。我們遵循 DP 的四個步驟:

步驟 1:分析最優解的結構 (Optimal Substructure)

核心思想:

任何一個 $A_i…A_j$ 的最優解,都必然是在某個 $k$ ( $i \le k < j$ ) 處切開,形成 $(A_i…A_k) \times (A_{k+1}…A_j)$。

並且,這個最優解所包含的「子問題解」$A_i…A_k$ 和 $A_{k+1}…A_j$ 也必須是它們各自的最優解。

(如果子問題不是最優,我們總能換成一個更優的子解,從而得到一個比原解更優的解,產生矛盾。)

步驟 2:建立遞迴解 (Recursive Solution)

  • 狀態定義:

    • m[i, j] = 計算矩陣鏈 $A_i…A_j$ 所需的最小純量乘法次數。

    • 我們的最終目標是 m[1, n]。

  • 邊界條件 (Base Case):

    • m[i, i] = 0 (單一矩陣,不需計算)
  • 狀態轉移方程:

    • 我們必須嘗試所有可能的切點 $k$ (從 $i$ 到 $j-1$)。

    • 成本 = (算左半邊) + (算右半邊) + (兩邊相乘)。

    • $m[i, j] = \min_{i \le k < j} {m[i, k] + m[k+1, j] + p_{i-1}p_kp_j}$

    • (其中 $p$ 是維度陣列,$A_i$ 的維度是 $p_{i-1} \times p_i$)

步驟 3:計算最優成本 (Computing Costs)

我們使用 Bottom-Up (由下而上) 的方式填表。

01-步驟 3 計算最優成本 (Computing Costs)

  • 演算法邏輯:

    1. 外層迴圈 l (鏈長度): 從 2 跑到 $n$。

    2. 中層迴圈 i (鏈起點): 從 1 跑到 $n-l+1$。

    3. 計算 j (鏈終點): j = i + l - 1。

    4. 內層迴圈 k (切點): 從 $i$ 跑到 $j-1$,用「狀態轉移方程」找最小值。

  • 和滑動視窗的關聯:

    • 當 l (鏈長度) 固定時,i 和 j 的迴圈就像一個「固定長度 l 的滑動視窗」,遍歷所有長度為 l 的子鏈。

    • 這是遍歷所有子問題的巧妙技巧,但演算法本質是 Interval DP,因為「長區間」的解依賴於「短區間」的解。

輸出方法

02-輸出方法

  • 如果 $i == j$ (只剩一個矩陣): 直接印出矩陣名字 (如 A1)。

  • 如果 $i < j$ (多個矩陣):

    • 它先印一個 (。

    • 接著,它去查 s[i, j] 得到最佳切點 $k$。

    • 它叫自己去遞迴印出「左半邊」 (從 $i$ 到 $k$)。

    • 再叫自己去遞迴印出「右半邊」 (從 $k+1$ 到 $j$)。

    • 最後,它印一個 )。

步驟 4:建構最佳解 (Constructing Solution)

  • 我們需要一個輔助表格 s[i, j]。

  • s[i, j] 儲存:在計算 m[i, j] 時,那個讓我們得到最小值的切點 $k$。

  • 回溯 (Backtracking):

    1. 從 s[1, n] 開始,得到 $k$。

    2. 這代表最終的括號是 $(A_1…A_k)(A_{k+1}…A_n)$。

    3. 遞迴地去 s[1, k] 和 s[k+1, n] 找下一層的括號。

4. ==範例演繹 (Walkthrough)==

計算 $A_1A_2A_3A_4$,維度 $p = [10, 100, 5, 50, 20]$ 這邊可以想像 $A_{1}$ 他是 $10 \times 100$,那 $A_2$ 是 $100 \times 5$,所以就會可以知道 $P_0=10,P_1=100$ 以此類推。

  • $A_1$: $10 \times 100$

  • $A_2$: $100 \times 5$

  • $A_3$: $5 \times 50$

  • $A_4$: $50 \times 20$

演算法會填滿以下兩個表格。填表的順序是由主對角線 (l=1) 開始,逐層往右上角 (l=4) 填,會先從區間長度 $l=2$ 開始,算出 $m[1,2]$ (成本 $5000$), $m[2,3]$ (成本 $25000$), $m[3,4]$ (成本 $5000$)。接著,區間長度 $l$ 變為 $3$,我們來算 $m[1,3]$ (也就是 $A_1A_2A_3$)。

這時有兩種切法 ( $k=1$ 或 $k=2$ ):

  1. $k=1$:切法是 $(A_1)(A_2A_3)$。

    • 成本 = m[1,1] (算 $A_1$) + m[2,3] (算 $A_2A_3$) + (兩者相乘的成本)

    • 成本 = $0 + 25000 + (p_0 \times p_1 \times p_3)$

    • 成本 = $0 + 25000 + (10 \times 100 \times 50) = 25000 + 50000 = 75000$

  2. $k=2$:切法是 $(A_1A_2)(A_3)$。

    • 成本 = m[1,2] (算 $A_1A_2$) + m[3,3] (算 $A_3$) + (兩者相乘的成本)

    • 成本 = $5000 + 0 + (p_0 \times p_2 \times p_3)$

    • 成本 = $5000 + 0 + (10 \times 5 \times 50) = 5000 + 2500 = 7500$

比較兩種切法:$75000$ vs $7500$,很明顯的 k=2 成本更低。

所以,演算法會更新:

  • m[1,3] = 7500

  • s[1,3] = 2 (記錄 $k=2$ 是最佳切點)

最終成本表 (m table)

m[i, j] = 計算 $A_i…A_j$ 的最小成本

i \ j 1 2 3 4
1 0 5000 7500 11000
2 - 0 25000 15000
3 - - 0 5000
4 - - - 0

最佳切點表 (s table)

s[i, j] = 計算 $A_i…A_j$ 時,得到最小成本的最佳切點 $k$

i \ j 1 2 3 4
1 - 1 2 2
2 - - 2 2
3 - - - 3
4 - - - -

最終結果

  • 最小成本: m[1, 4] = 11,000

  • 最佳加括號方式 (回溯 s 表):

    1. 看 s[1, 4] = 2 $\implies$ 切點 $k=2$ $\implies$ $((A_1A_2)(A_3A_4))$

    2. 看左邊 s[1, 2] = 1 $\implies$ 切點 $k=1$ $\implies$ $((A_1)(A_2))$

    3. 看右邊 s[3, 4] = 3 $\implies$ 切點 $k=3$ $\implies$ $((A_3)(A_4))$

最終結果

  • 最小成本: m[1, 4] = 11,000

  • 最佳加括號方式 (回溯 s 表):

    1. 看 s[1, 4] = 2 $\implies$ 切點 $k=2$ $\implies$ $((A_1A_2)(A_3A_4))$

    2. 看左邊 s[1, 2] = 1 $\implies$ 切點 $k=1$ $\implies$ $((A_1)(A_2))$

    3. 看右邊 s[3, 4] = 3 $\implies$ 切點 $k=3$ $\implies$ $((A_3)(A_4))$