1. 问题描述

给定 n 个矩阵 A_1, A_2, \dots, A_n​,其中矩阵 A_i​ 的维度为 p_{i-1} \times p_i(i=1,\dots,n),目标是确定矩阵连乘的计算顺序(即括号的插入方式),使得标量乘法的计算次数最小。注意矩阵乘法满足结合律,但不同的括号顺序可能导致不同的计算量。


2. 理论原理推导

2.1 最优子结构性质

假设我们要计算矩阵链 A_iA_{i+1}\cdots A_j​ 的最优计算顺序,设最优的括号分割点为 k(其中 i≤k<j),则必有:

其中:

  • m[i,j] 表示计算A_iA_{i+1}\cdots A_j 的最小计算代价(标量乘法次数)。

  • 转移方程(递推关系)为:

    m[i,j] = \min_{i\le k<j} \Big\{ m[i,k] + m[k+1,j] + p_{i-1} \times p_k \times p_j \Big\}.

这一公式体现了 最优子结构:如果一个问题的最优解包含了子问题的最优解,则可以通过递归地构造整个问题的最优解。

2.2 重叠子问题

对于不同的矩阵链区间 [i,j],可能会用到相同的子区间解 m[i,k] 或 m[k+1,j] 。直接递归求解会重复计算很多子问题,故采用动态规划将这些子问题解存储起来,避免重复计算。


3. 时间复杂度推导

设矩阵数量为 n :

  • 状态数: 动态规划表中有 \frac{n(n+1)}{2} 个子问题(只考虑 i ≤ j)。

  • 转移过程: 对于每个状态 m[i,j],需枚举所有可能的 k 值,最多 O(n) 次比较。

  • 总体时间复杂度: 因此总时间复杂度为

    O\left(\frac{n(n+1)}{2} \times n\right)=O(n^3).
  • 空间复杂度: DP 表需要 O(n^2) 空间。


4. 算法步骤

4.1 初始化

  • 构造一个二维数组m[1 \dots n][1 \dots n]用于存储各个子问题的最小计算代价。

  • 初始化对角线:对于 i=1,\dots,n,令 m[i,i]=0,因为单个矩阵没有乘法计算。

4.2 动态规划递推

  • 考虑链长 l 从 2 到 n :
    对于每个链长 l,遍历所有可能的起始位置 i(满足 1 \le i \le n-l+1),令 j=i+l-1。
    对于每个区间 [i,j]:

    • 初始化 m[i,j] = \infty。

    • 枚举分割点 k 从 i 到j-1,计算代价:

      q = m[i,k] + m[k+1,j] + p_{i-1} \times p_k \times p_j
    • 如果 q < m[i,j],则更新 m[i,j]=q 并记录相应的分割点 k(可选,用于重构最优括号方案)。

4.3 重构最优解(可选)

  • 如果需要输出最优的括号插入方案,可以维护一个辅助数组 s[i,j] 来记录使得 m[i,j]取得最小值时的分割点 k。

  • 使用递归方式或栈将 s 数组转化为具体的括号表示。


5. 示例代码

下面给出伪代码示例,展示如何填表计算最优矩阵链乘法代价:

# 输入:矩阵维度数组 p[0..n]
# 输出:最小计算代价 m[1][n],以及分割点 s[1][n]

def matrix_chain_order(p):
    n = len(p) - 1  # 矩阵个数
    # 初始化 m 和 s 两个二维数组
    m = [[0 if i == j else float('inf') for j in range(n+1)] for i in range(n+1)]
    s = [[0 for _ in range(n+1)] for _ in range(n+1)]
    
    # l 表示链长,从 2 到 n
    for l in range(2, n+1):
        for i in range(1, n - l + 2):
            j = i + l - 1
            for k in range(i, j):
                q = m[i][k] + m[k+1][j] + p[i-1] * p[k] * p[j]
                if q < m[i][j]:
                    m[i][j] = q
                    s[i][j] = k
    
    return m, s

def print_optimal_parens(s, i, j):
    if i == j:
        print(f"A{i}", end="")
    else:
        print("(", end="")
        print_optimal_parens(s, i, s[i][j])
        print_optimal_parens(s, s[i][j] + 1, j)
        print(")", end="")

# 示例调用
p = [30, 35, 15, 5, 10, 20, 25]
m, s = matrix_chain_order(p)
print("最小乘法次数为:", m[1][len(p)-1])
print("最优括号方案为: ", end="")
print_optimal_parens(s, 1, len(p)-1)

更多推荐