最长公共子序列(LCS)算法

给定两个字符串 s1 和 s2,找到最长公共子序列的长度。

最长公共子序列(LCS)算法
梯形图转SCL | 博途AI辅助编程文档 | AI模型价格对比 | AI工具导航 | ONNX模型库 | Vibe Coding教程 | PLC在线仿真器 | Tripo 3D | Meshy AI | ElevenLabs | KlingAI | ArtSpace | Phot.AI | InVideo

给定两个字符串 s1 和 s2,找到最长公共子序列的长度。如果没有公共子序列,返回 0。子序列是通过删除原始字符串中的 0 个或多个字符生成的字符串,不改变剩余字符的相对顺序。

例如,"ABC" 的子序列是 ""、"A"、"B"、"C"、"AB"、"AC"、"BC" 和 "ABC"。一般来说,长度为 n 的字符串有 2^n 个子序列。

示例:

输入: s1 = "ABC", s2 = "ACD" 
输出: 2 
解释: 两个字符串中存在的最长子序列是 "AC"。

输入: s1 = "AGGTAB", s2 = "GXTXAYB" 
输出: 4 
解释: 最长公共子序列是 "GTAB"。

输入: s1 = "ABC", s2 = "CBA" 
输出: 1 
解释: 有三个长度为 1 的最长公共子序列:"A"、"B" 和 "C"。

1、[朴素方法] 递归

成本:O(2^min(m, n)) 时间和 O(min(m, n)) 空间。

思路是比较 s1 和 s2 的最后一个字符。在比较字符串 s1 和 s2 时会出现两种情况:

  • 匹配:对剩余字符串(长度为 m-1 和 n-1 的字符串)进行递归调用,并将结果加 1。
  • 不匹配:进行两次递归调用。第一次对长度 m-1 和 n,第二次对 m 和 n-1。取两个结果的最大值。
  • 基本情况:如果任一字符串变为空,返回 0。

例如,考虑输入字符串 s1 = "ABX" 和 s2 = "ACX":

LCS("ABX", "ACX") = 1 + LCS("AB", "AC") [最后一个字符匹配]
LCS("AB", "AC") = max(LCS("A", "AC"), LCS("AB", "A")) [最后一个字符不匹配]
LCS("A", "AC") = max(LCS("", "AC"), LCS("A", "A")) = max(0, 1 + LCS("", "")) = 1
LCS("AB", "A") = max(LCS("A", "A"), LCS("AB", "")) = max(1 + LCS("", ""), 0) = 1

所以总体结果是 1 + 1 = 2

def lcsRec(s1, s2, m, n):
    # 基本情况:如果任一字符串为空,LCS 的长度为 0
    if m == 0 or n == 0:
        return 0

    # 如果两个子字符串的最后一个字符匹配
    if s1[m - 1] == s2[n - 1]:
        # 将此字符包含在 LCS 中,并对剩余子字符串进行递归
        return 1 + lcsRec(s1, s2, m - 1, n - 1)
    else:
        # 如果最后一个字符不匹配
        # 对两种情况进行递归:
        # 1. 排除 s1 的最后一个字符
        # 2. 排除 s2 的最后一个字符
        # 取这两个递归调用的最大值
        return max(lcsRec(s1, s2, m, n - 1), lcsRec(s1, s2, m - 1, n))

def lcs(s1, s2):
    m = len(s1)
    n = len(s2)
    return lcsRec(s1, s2, m, n)

if __name__ == "__main__":
    s1 = "AGGTAB"
    s2 = "GXTXAYB"
    print(lcs(s1, s2))

输出

4

2、[改进方法 1] 记忆化或自顶向下 DP

成本:O(m * n) 时间和 O(m * n) 空间

如果对字符串 "AXYT" 和 "AYZX" 使用上述递归方法,我们将得到如下的部分递归树。这里我们可以看到子问题 L("AXY", "AYZ") 被计算了多次。

如果考虑整棵树,会有许多这样的重叠子问题。因此我们可以使用记忆化或制表法来优化它。

  • 递归解决方案中有两个参数会变化,这些参数从 0 到 m 和 0 到 n。所以我们创建一个大小为 (m+1) x (n+1) 的二维数组。
  • 我们将此数组初始化为 -1,表示最初没有任何计算。
  • 现在我们修改递归解决方案,首先在此表中查找,如果值为 -1,则进行递归调用。
def lcsRec(s1, s2, m, n, memo):
    # 基本情况
    if m == 0 or n == 0:
        return 0

    # 已存在于记忆表中
    if memo[m][n] != -1:
        return memo[m][n]

    # 匹配
    if s1[m - 1] == s2[n - 1]:
        memo[m][n] = 1 + lcsRec(s1, s2, m - 1, n - 1, memo)
        return memo[m][n]

    # 不匹配
    memo[m][n] = max(lcsRec(s1, s2, m, n - 1, memo),
                     lcsRec(s1, s2, m - 1, n, memo))
    return memo[m][n]

def lcs(s1, s2):
    m = len(s1)
    n = len(s2)
    memo = [[-1] * (n + 1) for _ in range(m + 1)]
    return lcsRec(s1, s2, m, n, memo)

if __name__ == "__main__":
    s1 = "AGGTAB"
    s2 = "GXTXAYB"
    print(lcs(s1, s2))

输出

4

3、[改进方法 2] 自底向上 DP(制表法)

成本:O(m * n) 时间和 O(m * n) 空间。

递归解决方案中有两个参数会变化,这些参数从 0 到 m 和 0 到 n。所以我们创建一个大小为 (m+1) x (n+1) 的二维 dp 数组。
  • 首先填充 m 为 0 或 n 为 0 时的已知条目。
  • 然后使用递归公式填充剩余条目。

假设字符串是 S1 = "AXYT" 和 S2 = "AYZX",按照以下步骤:

def lcs(S1, S2):
    m = len(S1)
    n = len(S2)

    # 初始化大小为 (m+1)*(n+1) 的矩阵
    dp = [[0] * (n + 1) for x in range(m + 1)]

    # 自底向上构建 dp[m+1][n+1]
    for i in range(1, m + 1):
        for j in range(1, n + 1):
            if S1[i - 1] == S2[j - 1]:
                dp[i][j] = dp[i - 1][j - 1] + 1
            else:
                dp[i][j] = max(dp[i - 1][j], dp[i][j - 1])

    # dp[m][n] 包含 S1[0..m-1] 和 S2[0..n-1] 的 LCS 长度
    return dp[m][n]

if __name__ == "__main__":
    S1 = "AGGTAB"
    S2 = "GXTXAYB"
    print(lcs(S1, S2))

输出

4

4、[期望方法] 空间优化 - 单数组

成本:O(m * n) 时间和 O(n) 空间。

上述简单实现中的一个重要观察是,在外循环的每次迭代中,我们只需要前一行所有列的值。因此没有必要在 dp 矩阵中存储所有前面的行。

一种优化空间只存储前一行的方法是,我们也可以通过使用临时变量 prev 来避免存储前一行。

1D DP 表条目 dp[j] 表示更新前 dp[i-1][j](前一行的值)。在计算过程中,dp[j] 被更新以表示当前行值 dp[i][j],并为下一次迭代更新 prev。

递推关系变为:

  • 如果字符 s1[i-1] 和 s2[j-1] 匹配,dp[j] = 1 + prev。这里,prev 是存储对角线值 (dp[i-1][j-1]) 的临时变量。
  • 如果字符不匹配,dp[j] = max(dp[j-1], dp[j])。这里 dp[j] 表示更新前的 dp[i-1][j],dp[j-1] 表示 dp[i][j-1] 的值。计算当前值后,我们将 prev 更新为 dp[j] 的旧值,用于下一列。
def lcs(s1, s2):
    m = len(s1)
    n = len(s2)

    # dp 数组初始化为全零
    # 此数组存储当前行的 LCS 值。
    # dp[j] 表示 s1[0..i] 和 s2[0..j] 的 LCS
    dp = [0] * (n + 1)

    # i 和 j 分别表示 s1 和 s2 的长度
    for i in range(1, m + 1):

        # prev 存储前一行和前一列 (i-1), (j-1) 的值
        # 用于在更新 dp[j] 时跟踪 LCS[i-1][j-1]
        prev = dp[0]

        for j in range(1, n + 1):

            # temp 临时存储更新前的当前 dp[j]
            temp = dp[j]

            if s1[i - 1] == s2[j - 1]:
                # 如果字符匹配,从前一行和前一列的值加 1
                dp[j] = 1 + prev
            else:
                # 否则,取左边 (dp[j-1]) 和上方 (dp[j]) 值的最大值
                dp[j] = max(dp[j - 1], dp[j])

            # 为下一次迭代更新 prev
            prev = temp

    # 数组的最后一个元素包含 LCS 的长度
    return dp[n]

if __name__ == "__main__":
    s1 = "AGGTAB"
    s2 = "GXTXAYB"
    print(lcs(s1, s2))

输出

4

原文链接: Longest Common Subsequence (LCS)

汇智网翻译整理,转载请标明出处