题解:P17236 『STA - R10』刷墙墙刷
黄题做了
这里字符串下标从
Subtask 1&2
我们先思考满足什么情况下,字符串
-
- 存在任意满足
3 \le i \le n 的正整数,使b_i=b_{i-2} 。
于是枚举所有的
:::success[参考代码]
#include<iostream>
#include<string>
using namespace std;
const int mod = 998244353;
string a, b, s; int n, ans = 0;
void dfs(int x){
if(x == n){
if(a == s){
ans++;
return;
}
for(int i = 2; i < s.size(); i++){
if(s[i] == s[i-2]){
ans++;
return;
}
}
return;
}
if(b[x] == '?'){
for(char i = 'a'; i <= 'z'; i++){
s.push_back(i);
dfs(x+1);
s.pop_back();
}
}
else{
s.push_back(b[x]);
dfs(x+1);
s.pop_back();
}
}
int main(){
cin >> n >> a >> b;
dfs(0);
cout << ans;
return 0;
}
:::
Subtask4
前面说明的两种情况,第一种显然最多贡献
- 存在任意满足
1 \le i \le n 的正整数,使a_i \ne b_i 且b_i= ?。 - 存在任意满足
3 \le i \le n 的正整数,使a_i=a_{i-2} 如果没有,则可直接将答案增加1 。
考虑 dp,设
- 当
s_1=s_3 是,有f_{i,s_2,s_3,1} \leftarrow f_{i,s_2,s_3,1}+f_{i,s_1,s_2,0} - 当
s_1 \ne s_3 是,有f_{i,s_2,s_3,0} \leftarrow f_{i,s_2,s_3,0}+f_{i,s_1,s_2,0} - 所有情况都有
f_{i,s_2,s_3,1} \leftarrow f_{i,s_2,s_3,1}+f_{i,s_1,s_2,1}
朴素枚举转移,设字符集为
:::success[参考代码]
#include<iostream>
#include<string>
using namespace std;
const int N = 1e6 + 10, M = 30, mod = 998244353;
string a, b, s;
int n, ans = 0, f[2][M][M][2];
bool sub2 = 1, fg;
void dfs(int x){
if(x == n){
if(a == s){
ans++;
return;
}
for(int i = 2; i < s.size(); i++){
if(s[i] == s[i-2]){
ans++;
return;
}
}
return;
}
if(b[x] == '?'){
for(char i = 'a'; i <= 'z'; i++){
s.push_back(i);
dfs(x+1);
s.pop_back();
}
}
else{
s.push_back(b[x]);
dfs(x+1);
s.pop_back();
}
}
int main(){
cin >> n >> a >> b; fg = 1;
for(int i = 0; i < n; i++){
if(b[i] == '?') sub2 = 0;
if(i >= 2 && a[i] == a[i-2]) fg = 0;
if(c1[i] != c2[i] && c2[i] != '?') fg = 0;
}
bool sub1 = (n <= 5);
if(sub1 || sub2){
dfs(0);
cout << ans;
}
else{
ans += fg;
for(char i = 'a'; i <= 'z'; i++)
for(char j = 'a'; j <= 'z'; j++)
if((b[0] == '?' || b[0] == i) && (b[1] == '?' || b[1] == j)) f[0][i-'a'][j-'a'][0] = 1;
for(int i = 3; i <= n; i++){
for(char s1 = 'a'; s1 <= 'z'; s1++){
if(b[i-3] != '?' && b[i-3] != s1) continue;
for(char s2 = 'a'; s2 <= 'z'; s2++){
if(b[i-2] != '?' && b[i-2] != s2) continue;
for(char s3 = 'a'; s3 <= 'z'; s3++){
if(b[i-1] != '?' && b[i-1] != s3) continue;
if(s1 == s3) f[i&1][s2-'a'][s3-'a'][1] = (f[i&1][s2-'a'][s3-'a'][1] + f[(i&1)^1][s1-'a'][s2-'a'][0]) % mod;
else f[i&1][s2-'a'][s3-'a'][0] = (f[i&1][s2-'a'][s3-'a'][0] + f[(i&1)^1][s1-'a'][s2-'a'][0]) % mod;
f[i&1][s2-'a'][s3-'a'][1] = (f[i&1][s2-'a'][s3-'a'][1] + f[(i&1)^1][s1-'a'][s2-'a'][1]) % mod;
}
}
}
for(char s1 = 'a'; s1 <= 'z'; s1++){
for(char s2 = 'a'; s2 <= 'z'; s2++){
f[(i&1^1)][s1-'a'][s2-'a'][0] = 0;
f[(i&1^1)][s1-'a'][s2-'a'][1] = 0;
}
}
}
for(char s1 = 'a'; s1 <= 'z'; s1++){
for(char s2 = 'a'; s2 <= 'z'; s2++){
ans = (ans + f[n&1][s1-'a'][s2-'a'][1]) % mod;
}
}
cout << ans;
}
return 0;
}
:::
Subtask3
令字符集为
可以发现,
-
g_{i,s_2,0}=\sum_{s_1 \in S}=f_{i,s1,s2,0} -
g_{i,s_2,1}=\sum_{s_1 \in S}=f_{i,s1,s2,1}
则转移式子可以转化为:
-
f_{i,s_1,s_2,0}=g_{i,s_2,0}-f_{i,s_2,s_1,0} -
f_{i,s_1,s_2,1}=g_{i,s_2,1}+f_{i,s_2,s_1,0}
进行前缀和优化,复杂度
:::success[参考代码]
#include<iostream>
#include<string>
using namespace std;
const int N = 1e6 + 10, M = 30, mod = 998244353;
string a, b, s;
int n, ans = 0, f[2][M][M][2], g[M][2];
bool fg;
int main(){
cin >> n >> a >> b; fg = 1;
for(int i = 0; i < n; i++){
if(i >= 2 && a[i] == a[i-2]) fg = 0;
if(a[i] != b[i] && b[i] != '?') fg = 0;
}
ans += fg;
for(char i = 'a'; i <= 'z'; i++)
for(char j = 'a'; j <= 'z'; j++){
if((b[0] == '?' || b[0] == i) && (b[1] == '?' || b[1] == j)) f[0][i-'a'][j-'a'][0] = 1;
g[j-'a'][0] = (g[j-'a'][0] + f[0][i-'a'][j-'a'][0]) % mod;
}
for(int i = 3; i <= n; i++){
// for(char s1 = 'a'; s1 <= 'z'; s1++){
// if(b[i-3] != '?' && b[i-3] != s1) continue;
// for(char s2 = 'a'; s2 <= 'z'; s2++){
// if(b[i-2] != '?' && b[i-2] != s2) continue;
// for(char s3 = 'a'; s3 <= 'z'; s3++){
// if(b[i-1] != '?' && b[i-1] != s3) continue;
// f[i&1][s2-'a'][s3-'a'][0] = (f[i&1][s2-'a'][s3-'a'][0] + f[(i&1)^1][s1-'a'][s2-'a'][0]) % mod;
// }
// }
// }
for(char s1 = 'a'; s1 <= 'z'; s1++){
if(b[i-2] != '?' && b[i-2] != s1) continue;
for(char s2 = 'a'; s2 <= 'z'; s2++){
if(b[i-1] != '?' && b[i-1] != s2) continue;
f[i&1][s1-'a'][s2-'a'][1] = (f[i&1][s1-'a'][s2-'a'][1] + f[(i&1)^1][s2-'a'][s1-'a'][0]) % mod;
f[i&1][s1-'a'][s2-'a'][1] = (f[i&1][s1-'a'][s2-'a'][1] + g[s1-'a'][1]) % mod;
f[i&1][s1-'a'][s2-'a'][0] = (f[i&1][s1-'a'][s2-'a'][0] - f[(i&1)^1][s2-'a'][s1-'a'][0] + mod) % mod;
f[i&1][s1-'a'][s2-'a'][0] = (f[i&1][s1-'a'][s2-'a'][0] + g[s1-'a'][0]) % mod;
}
}
for(char s1 = 'a'; s1 <= 'z'; s1++) g[s1-'a'][0] = g[s1-'a'][1] = 0;
for(char s1 = 'a'; s1 <= 'z'; s1++){
for(char s2 = 'a'; s2 <= 'z'; s2++){
f[(i&1^1)][s1-'a'][s2-'a'][0] = 0;
f[(i&1^1)][s1-'a'][s2-'a'][1] = 0;
g[s2-'a'][1] = (g[s2-'a'][1] + f[i&1][s1-'a'][s2-'a'][1]) % mod;
g[s2-'a'][0] = (g[s2-'a'][0] + f[i&1][s1-'a'][s2-'a'][0]) % mod;
}
}
}
for(char s1 = 'a'; s1 <= 'z'; s1++){
for(char s2 = 'a'; s2 <= 'z'; s2++){
ans = (ans + f[n&1][s1-'a'][s2-'a'][1]) % mod;
}
}
cout << ans;
return 0;
}
:::
Subtask5
观察到方案数可以将奇数和偶数分开算,设
- 存在任意满足
2 \le i \le |s_0| 的正整数,使s_{0,i}=s_{0,i-1} 。 - 存在任意满足
2 \le i \le |s_1| 的正整数,使s_{1,i}=s_{1,i-1} 。
dp 方程就与上面同理,然后可以求出:
于是,答案为:
那么上面的代码就报废了,重新写完,朴素枚举复杂度为
:::success[参考代码]
#include<iostream>
#include<string>
using namespace std;
typedef long long ll;
const int N = 1e6 + 10, M = 30, mod = 998244353;
int n, f[N][M][2], ans;
string c1, c2, a1, b1, a2, b2;
bool fg;
pair<int, int> solve(string & a, string & b){
n = a.size();
if (b.empty()) return {1, 0};
for(int i = 1; i <= n; i++){
for(char j = 'a'; j <= 'z'; j++) f[i][j-'a'][0] = f[i][j-'a'][1] = 0;
}
for(char i = 'a'; i <= 'z'; i++){
f[1][i-'a'][0] = 1;
}
// cout << a << ' ' << b << '\n';
for(int i = 2; i <= n; i++){
for(char s1 = 'a'; s1 <= 'z'; s1++){
if(b[i-2] != '?' && b[i-2] != s1) continue;
for(char s2 = 'a'; s2 <= 'z'; s2++){
if(b[i-1] != '?' && b[i-1] != s2) continue;
// cout << i << ' ' << s1 << ' ' << s2 << ' ' << f[i][s1-'a'][0] << ' ' << f[i-1][s2-'a'][0] << '\n';
if(s1 == s2) f[i][s2-'a'][1] = (f[i][s2-'a'][1] + f[i-1][s1-'a'][0]) % mod;
else f[i][s2-'a'][0] = (f[i][s2-'a'][0] + f[i-1][s1-'a'][0]) % mod;
f[i][s2-'a'][1] = (f[i][s2-'a'][1] + f[i-1][s1-'a'][1]) % mod;
}
}
}
int res1 = 0, res2 = 0;
for(char i = 'a'; i <= 'z'; i++){
res1 = (res1 + f[n][i-'a'][0]) % mod;
res2 = (res2 + f[n][i-'a'][1]) % mod;
}
return {res1, res2};
}
int main(){
cin >> n >> c1 >> c2; fg = 1;
for(int i = 0; i < n; i++){
if(i&1){
a1.push_back(c1[i]);
b1.push_back(c2[i]);
}
else{
a2.push_back(c1[i]);
b2.push_back(c2[i]);
}
}
for(int i = 0; i < n; i++){
if(i >= 2 && c1[i] == c1[i-2]) fg = 0;
if(c1[i] != c2[i] && c2[i] != '?') fg = 0;
}
pair<int, int> p1 = solve(a1, b1);
pair<int, int> p2 = solve(a2, b2);
ans = fg;
ans = (ans + (ll)p1.first*p2.second) % mod;
ans = (ans + (ll)p1.second*p2.second) % mod;
ans = (ans + (ll)p1.second*p2.first) % mod;
// cout << p1.first << ' ' << p1.second << ' ' << p2.first << ' ' << p2.second << '\n';
cout << ans;
return 0;
}
:::
前缀和优化后复杂度为
:::success[参考代码]
#include<iostream>
#include<string>
using namespace std;
typedef long long ll;
const int N = 1e6 + 10, M = 30, mod = 998244353;
int n, f[N][M][2], g[2], ans;
string c1, c2, a1, b1, a2, b2;
bool fg;
pair<int, int> solve(string & a, string & b){
n = a.size(); g[0] = g[1] = 0;
if (b.empty()) return {1, 0};
for(int i = 1; i <= n; i++){
for(char j = 'a'; j <= 'z'; j++) f[i][j-'a'][0] = f[i][j-'a'][1] = 0;
}
for(char i = 'a'; i <= 'z'; i++){
if(b[0] != '?' && b[0] != i) continue;
f[1][i-'a'][0] = 1; g[0]++;
}
// cout << a << ' ' << b << '\n';
for(int i = 2; i <= n; i++){
// for(char s1 = 'a'; s1 <= 'z'; s1++){
// if(b[i-2] != '?' && b[i-2] != s1) continue;
// for(char s2 = 'a'; s2 <= 'z'; s2++){
// if(b[i-1] != '?' && b[i-1] != s2) continue;
//// cout << i << ' ' << s1 << ' ' << s2 << ' ' << f[i][s1-'a'][0] << ' ' << f[i-1][s2-'a'][0] << '\n';
// if(s1 == s2) f[i][s2-'a'][1] = (f[i][s2-'a'][1] + f[i-1][s1-'a'][0]) % mod;
// else f[i][s2-'a'][0] = (f[i][s2-'a'][0] + f[i-1][s1-'a'][0]) % mod;
// f[i][s2-'a'][1] = (f[i][s2-'a'][1] + f[i-1][s1-'a'][1]) % mod;
// }
// }
for(char s1 = 'a'; s1 <= 'z'; s1++){
if(b[i-1] != '?' && b[i-1] != s1) continue;
f[i][s1-'a'][1] = (f[i][s1-'a'][1] + f[i-1][s1-'a'][0]) % mod;
f[i][s1-'a'][1] = (f[i][s1-'a'][1] + g[1]) % mod;
f[i][s1-'a'][0] = (f[i][s1-'a'][0] - f[i-1][s1-'a'][0] + mod) % mod;
f[i][s1-'a'][0] = (f[i][s1-'a'][0] + g[0]) % mod;
}
g[0] = g[1] = 0;
for(char s1 = 'a'; s1 <= 'z'; s1++){
g[0] = (g[0] + f[i][s1-'a'][0]) % mod;
g[1] = (g[1] + f[i][s1-'a'][1]) % mod;
}
}
int res1 = 0, res2 = 0;
for(char i = 'a'; i <= 'z'; i++){
res1 = (res1 + f[n][i-'a'][0]) % mod;
res2 = (res2 + f[n][i-'a'][1]) % mod;
}
return {res1, res2};
}
int main(){
cin >> n >> c1 >> c2; fg = 1;
for(int i = 0; i < n; i++){
if(i&1){
a1.push_back(c1[i]);
b1.push_back(c2[i]);
}
else{
a2.push_back(c1[i]);
b2.push_back(c2[i]);
}
}
for(int i = 0; i < n; i++){
if(i >= 2 && c1[i] == c1[i-2]) fg = 0;
if(c1[i] != c2[i] && c2[i] != '?') fg = 0;
}
pair<int, int> p1 = solve(a1, b1);
pair<int, int> p2 = solve(a2, b2);
ans = fg;
ans = (ans + (ll)p1.first*p2.second) % mod;
ans = (ans + (ll)p1.second*p2.second) % mod;
ans = (ans + (ll)p1.second*p2.first) % mod;
// cout << p1.first << ' ' << p1.second << ' ' << p2.first << ' ' << p2.second << '\n';
cout << ans;
return 0;
}
:::