题解:P17236 『STA - R10』刷墙墙刷

· · 题解

黄题做了 3 小时,我怎么这么菜。

这里字符串下标从 1 开始。

Subtask 1&2

我们先思考满足什么情况下,字符串 a 可以通过题目中的操作转化为一个字符串 s,一共两种情况:

于是枚举所有的 b,检验 b 是否满足以上两个条件即可,计问号个数为 x,时间复杂度 \mathcal{O}(26^x+n),可以同时通过 Subtask 1&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 个解,这个解有可能在第二种也贡献一次,所以需要判断 a 是否满足:

考虑 dp,设 f_{i,j,k,0/1} 表示在第 i 个位置,最后一个字母为 j,倒数第二个字母为 k,是否已经满足某个 s_i=s_{i-2},这样 i=1 时会无意义,所以可以预处理出 i=2,再从 i=3 开始进行转移,枚举 s_1,s_2,s_3,范围是字符集,共三种情况:

朴素枚举转移,设字符集为 V,复杂度 \mathcal{O}(nV^3),加数据点分治可以过掉 Subtask1&2&4

:::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

令字符集为 S

可以发现,s_1 可以不枚举,设:

则转移式子可以转化为:

进行前缀和优化,复杂度 \mathcal{O}(nV^2),可以直接过掉 Subtask1&2&3&4

:::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

观察到方案数可以将奇数和偶数分开算,设 s_0,s_1 分别为 s 偶数/奇数编号的字符,则条件可以变为满足以下两个条件之一:

dp 方程就与上面同理,然后可以求出:

于是,答案为:

c_{0,1}c_{1,0}+c_{0,0}c_{1,1}+c_{0,1}c_{1,1}

那么上面的代码就报废了,重新写完,朴素枚举复杂度为 \mathcal{O}(nV^2)

:::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;
}

:::

前缀和优化后复杂度为 \mathcal{O}(nV)

:::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;
}

:::