如何高效统计字符串s中大于字符串t的子序列数量?
问题:统计字符串s中大于t的子序列数量
给定两个字符串s(长度m)和t(长度n),统计s中所有大于t的子序列数量。
子序列p大于q的判定条件:
- a) 在p和q第一个不同的位置i处,p[i] > q[i];
- b) p的长度|p|大于q的长度|q|,且q是p的前缀。
示例
s = "bab",t = "ab"
结果 = 5
解释:
符合条件的s的子序列为:"b"、"ba"、"bb"、"bab"、"b"
约束条件
- s的长度范围为1到10^5
- t的长度范围为1到100,t的长度也可能大于s的长度
原递归解法(时间复杂度O(2^m * m))
public class Main { private static final int MOD = 1_000_000_007; private static void subsequence(String s, int index, String current, List<String> subsequences) { if (index == s.length()) { if (!current.isEmpty()) { subsequences.add(current); } return; } subsequence(s, index + 1, current, subsequences); subsequence(s, index + 1, current + s.charAt(index), subsequences); } private static boolean isGreater(String s1, String t) { int len1 = s1.length(); int len2 = t.length(); for (int i = 0; i < Math.min(len1, len2); i++) { if (s1.charAt(i) > t.charAt(i)) { return true; } else if (s1.charAt(i) < t.charAt(i)) { return false; } } return len1 > len2; } public static int solve(String s, String t) { List<String> subsequences = new ArrayList<>(); subsequence(s, 0, "", subsequences); int count = 0; for (String e : subsequences) { if (isGreater(e, t)) { count = (count + 1) % MOD; } } return count; } public static void main(String[] args) { System.out.println(solve("aba", "ab")); // Expected: 3 System.out.println(solve("bab", "ab")); // Expected: 5 System.out.println(solve("wrrmkhds", "bebbjvcgzlwtbvasphvm")); // Expected: 255 System.out.println(solve("o", "h")); // Expected: 1 } }
原解法枚举所有子序列,当s长度为1e5时,2^1e5的复杂度完全不可行,必须用动态规划优化。
优化解法(时间复杂度O(m*n))
思路分析
我们将符合条件的子序列分为两类计算:
- 满足条件a的子序列:遍历s的每个字符,维护已匹配t前缀的子序列数目,若当前字符大于t对应前缀位置的字符,则所有已匹配该前缀的子序列加上当前字符及后续任意组合都符合条件。
- 满足条件b的子序列:所有以t为前缀且长度大于n的子序列,即匹配完t的全部字符后,后续至少选一个字符的组合。
具体步骤:
- 预计算2的幂次数组
pow2,快速获取剩余字符的组合数。 - 用
dp[j]记录已匹配t前j个字符的子序列数目,sumDp[j]记录匹配到t前j个字符时,后续字符组合数的总和。 - 遍历s的每个字符,先计算满足条件a的贡献,再倒序更新匹配状态(避免覆盖未处理的状态)。
- 最后加上满足条件b的子序列数目。
优化代码
public class Main { private static final int MOD = 1_000_000_007; public static int solve(String s, String t) { int m = s.length(); int n = t.length(); // 预计算2的幂次数组,pow2[i] = 2^i mod MOD long[] pow2 = new long[m + 1]; pow2[0] = 1; for (int i = 1; i <= m; i++) { pow2[i] = (pow2[i - 1] * 2) % MOD; } long[] dp = new long[n + 1]; dp[0] = 1; // 初始状态:匹配0个字符的子序列数目为1(空序列) long[] sumDp = new long[n + 1]; // 记录匹配到j个字符时,后续组合数的总和 long ans = 0; for (int i = 0; i < m; i++) { char c = s.charAt(i); // 计算满足条件a的贡献:当前字符大于t的第j个字符时,累加对应组合数 for (int j = 0; j < n; j++) { if (c > t.charAt(j)) { ans = (ans + dp[j] * pow2[m - i - 1]) % MOD; } } // 倒序更新匹配状态,避免覆盖未处理的dp[j] for (int j = n - 1; j >= 0; j--) { if (c == t.charAt(j)) { sumDp[j + 1] = (sumDp[j + 1] + dp[j] * pow2[m - i - 1]) % MOD; dp[j + 1] = (dp[j + 1] + dp[j]) % MOD; } } } // 计算满足条件b的数目:所有匹配完t的子序列的后续组合数总和,减去等于t的子序列数目 long caseB = (sumDp[n] - dp[n] + MOD) % MOD; ans = (ans + caseB) % MOD; return (int) ans; } public static void main(String[] args) { System.out.println(solve("aba", "ab")); // Expected: 3 System.out.println(solve("bab", "ab")); // Expected: 5 System.out.println(solve("wrrmkhds", "bebbjvcgzlwtbvasphvm")); // Expected: 255 System.out.println(solve("o", "h")); // Expected: 1 } }
内容的提问来源于stack exchange,提问作者CodeCrusader
相关产品推荐
相关产品推荐

