Interviewbit Stringoholics问题:排查代码错误(已知KMP解法)
DSA问题代码错误排查
问题背景
给定仅由'a'和'b'组成的字符串数组A,每个字符串执行以下操作:
- 时刻1:循环右移1个字符;
- 时刻2:对新字符串循环右移2个字符;
- 时刻3:对新字符串循环右移3个字符;
- 时刻i:对新字符串循环右移
i % 字符串长度个字符。
示例:字符串"abaa"
时刻1:变为"baaa";时刻2:变为"aaba";时刻3:变为"aaab";时刻4:仍为"aaab";时刻5:变为"aaba"
字符串变回原状态后流程重置:若需t时间回归原状态,则t+1时刻再次右移1个字符,2t时刻回归原状态。需找出最小时刻,使最多字符串同时处于原状态,结果对10^9+7取模。
约束与示例
- 约束:1 ≤ |A| ≤ 100
- 示例1:输入
A=['a','ababa','aba'],输出4; - 示例2:输入
A=['a','aa'],输出1。
现有代码
long long int gcd(int a, int b){ if(b==0) return a; return gcd(b,a%b); } long long int lcm(int a, int b){ return a*b/(gcd(a,b)); } int check(string a){ string b=a; int ans=0; int n=a.size(); for(int i=1;i<INT_MAX;i++){ a=a+a.substr(0,i%n); a.erase(a.begin(),a.begin()+i%n); cout<<a<<" "; ans++; if(a==b) return ans; } return n; } int Solution::solve(vector<string> &a) { int n=a.size(); int maxi=1000000009; long long int ans=1; vector<int>k; for(int i=0;i<n;i++){ k.push_back(check(a[i])); //cout<<k[i]<<" "; } for(int i=0;i<k.size();i++){ ans=lcm(ans,k[i]); //cout<<k[i]<<" "; } return (ans)%maxi; }
代码错误分析与修正建议
1. LCM计算的整数溢出问题
当前lcm函数参数为int类型,当两个参数取值较大时(例如超过46340,因为46340^2接近int类型最大值),a*b会直接溢出int范围,导致计算结果错误。同时,先乘后除的方式进一步增加了溢出风险。
修正方案:
将参数改为long long类型,且先除以最大公约数再相乘,减少溢出概率:
long long int gcd(long long int a, long long int b){ if(b==0) return a; return gcd(b,a%b); } long long int lcm(long long int a, long long int b){ if(a == 0 || b == 0) return 0; return a / gcd(a, b) * b; }
2. 模拟操作的效率与死循环风险
check函数通过循环模拟每一步操作直到字符串回归原状态,当回归周期较大时(例如1e6级别),循环会持续极长时间,甚至因INT_MAX限制导致死循环或超时。此外,模拟过程中频繁的字符串拼接与删除操作也会大幅降低效率。
修正方案:
利用KMP算法计算字符串的最小周期,再通过数学推导直接求解回归周期,无需模拟:
- 第一步:用KMP求字符串的部分匹配表(next数组),推导最小周期
d:vector<int> computeNext(string s) { int n = s.size(); vector<int> next(n, 0); for(int i=1;i<n;i++){ int j = next[i-1]; while(j>0 && s[i]!=s[j]){ j = next[j-1]; } if(s[i]==s[j]){ j++; } next[i] = j; } return next; } int getMinPeriod(string s) { int n = s.size(); vector<int> next = computeNext(s); int d = n - next.back(); return n % d == 0 ? d : n; } - 第二步:根据最小周期
d,求解满足累计位移和模d为0的最小t:long long findMinT(int n, int d) { // 累计位移和公式:sum = q*n*(n-1)/2 + r*(r+1)/2,t = q*n + r for(long long t=1;t<=1e18;t++){ long long q = t / n; long long r = t % n; long long sum = q * (long long)n * (n-1)/2 + r*(r+1)/2; if(sum % d == 0){ return t; } } return n; }
3. 流程逻辑的隐性错误
当前代码默认所有字符串的回归周期的最小公倍数就是答案,但需注意:若字符串本身是全相同字符(如"aaaa"),其回归周期为1,任何时刻都处于原状态,需特殊处理(上述修正代码已覆盖该场景)。
修正后的完整核心代码
#include <vector> #include <string> using namespace std; const int MOD = 1e9+7; long long gcd(long long a, long long b){ if(b==0) return a; return gcd(b,a%b); } long long lcm(long long a, long long b){ if(a == 0 || b == 0) return 0; return a / gcd(a, b) * b; } vector<int> computeNext(string s) { int n = s.size(); vector<int> next(n, 0); for(int i=1;i<n;i++){ int j = next[i-1]; while(j>0 && s[i]!=s[j]){ j = next[j-1]; } if(s[i]==s[j]){ j++; } next[i] = j; } return next; } int getMinPeriod(string s) { int n = s.size(); vector<int> next = computeNext(s); int d = n - next.back(); return n % d == 0 ? d : n; } long long findMinT(int n, int d) { for(long long t=1;t<=1e18;t++){ long long q = t / n; long long r = t % n; long long sum = q * (long long)n * (n-1)/2 + r*(r+1)/2; if(sum % d == 0){ return t; } } return n; } int Solution::solve(vector<string> &a) { vector<long long> periods; for(string s : a){ int n = s.size(); if(n == 0){ periods.push_back(1); continue; } int d = getMinPeriod(s); long long t = findMinT(n, d); periods.push_back(t); } long long ans = 1; for(long long p : periods){ ans = lcm(ans, p); ans %= MOD; } return ans % MOD; }
内容的提问来源于stack exchange,提问作者Premlani Shailesh
相关产品推荐
相关产品推荐

