[AGC045C] Range Set

· · 题解

或许更好的阅读体验。

思路:

这种算最终能形成多少本质不同的东西,都是考虑倒着做,即给定一个字符串,判断它是否能被操作给表示出来。

这题也是一样,考虑怎么判定字符串 s 能被操作出来,你发现正着完全无法做,那么倒着考虑,知道最后一次操作后,这些位置在之前可以被任意赋值,这样限制很少了。

具体的,显然最后一次一定是连着的长为 a0 或者长为 b1,随意找到一个这样区间 [l, r],将里面字符改为 ? 表示通配符,倒数第二次也是一样找长 a, b 的连续段(算上通配符),然后继续改为 ?

所以本质上,判定 s 是否合法,当且仅当,可以通过每次选择一个长为 a0 或者长为 b1(也可以包含 ? 通配符)段,将其改成 ?,使得最终全是 ?

先钦定 a \ge b,考虑下这个的最优策略是什么?显然,对于任何长度 \ge b 的全 1 段一定可以变成全 ?,并且操作这个是不劣的,然后考虑剩下的 10 怎么办:

所以有解当且仅当把长度 \ge b 的全 1 段覆盖为 ? 后,存在一个长为 a0? 的连续段,这是充要的,所以可以直接对这个计数;但是存在还是不太好做,显然容斥一下算一下不存在长度 \ge a 的即可。

然后考虑怎么计数?显然对于长度 \ge b 的全 1 段,可以看作是 0,于是转移的时候,要么是加上一个长度 <b 的全 1 段,要么是加上一个 <a 的“全 0 段”,于是为了辅助,定义 f_i 表示长度为 i 的满足所有连续 1 段长度 \ge b01 序列数(当然也可以不存在 1)。

于是你的结构一定是 <a 的“全 0 段”与 <b 的全 1 段交替拼起来。

然后再定义 dp_{i, 0/1} 表示以 i 结尾不合法的本质不同字符串个数,且结尾是 0/1,有转移:

dp_{i, 0} = \sum_{i - j + 1 < a} f_{i - j - (j > 1)} dp_{j - 1, 1} dp_{i, 1} = \sum_{i - j + 1 < b} dp_{j - 1, 0}

最后答案怎么算,你发现直接 2^n - dp_{n, 0} - dp_{n ,1} 是有问题的,具体的,你发现上面转移实际上是有问题的;对于一个“全 0 段”最后一段结尾不一定是 0,还可能是 1;那显然只有最后一个“全 0 段”结尾时末尾可以是 1,因为如果在中间结尾是 1 后面还拼了个 <b 的全 1 段是不合法的。

再定义对于一个长度为 i 的以 0/1 结尾的“全 0 段”是 g_{i, 0/1},显然:

g_{i, 0} = f_{i - 1} g_{i, 1}= \sum_{j = 1}^{i - b + 1} f_{\max(0, j - 2)}

顺便转移下 f,显然有:

f_i = g_{i, 0} + g_{i, 1}

所以实际上答案是:

2^n - dp_{n, 1} - dp_{n, 0} - \sum_{n - i + 1 < a} g_{n - i + 1 - [i > 1], 1} dp_{i - 1, 1}

时间复杂度为 O(n^2)

完整代码:

#include<bits/stdc++.h>
#define fi first
#define se second
#define lowbit(x) (x) & (-(x))
#define popcnt(x) __builtin_popcount(x)
using namespace std;
typedef unsigned long long ull;
typedef long long ll;
const int N = 5e3 + 10, mod = 1e9 + 7;
inline ll read(){
    ll x = 0, f = 1;
    char c = getchar();
    while(c < '0' || c > '9'){
        if(c == '-')
          f = -1;
        c = getchar();
    }
    while(c >= '0' && c <= '9'){
        x = (x << 1) + (x << 3) + (c ^ 48);
        c = getchar();
    }
    return x * f;
}
inline void write(ll x){
    if(x < 0){
        putchar('-');
        x = -x;
    }
    if(x > 9)
      write(x / 10);
    putchar(x % 10 + '0');
}
inline void getadd(int &x, int y){
    x = (x + y >= mod) ? (x + y - mod) : (x + y);
}
inline int add(int x, int y){
    return (x + y >= mod) ? (x + y - mod) : (x + y);
}
inline void getdec(int &x, int y){
    x = (x < y) ? (x - y + mod) : (x - y);
}
inline int dec(int x, int y){
    return (x < y) ? (x - y + mod) : (x - y);
}
int n, a, b, ans;
int f[N], g[N][2], dp[N][N];
int main(){
    n = read(), a = read(), b = read();
    if(a < b)
      swap(a, b);
    f[0] = ans = 1;
    for(int i = 1; i <= n; ++i){
        ans = ans * 2ll % mod;
        g[i][0] = f[i - 1];
        for(int j = 1; j <= i - b + 1; ++j)
          getadd(g[i][1], f[max(0, j - 2)]);
        f[i] = add(g[i][0], g[i][1]);
        // cerr << f[i] << '\n';
    }
    dp[0][0] = dp[0][1] = 1;
    for(int i = 1; i <= n; ++i){
        for(int j = 1; j <= i; ++j){
            if(i - j + 1 < a)
              getadd(dp[i][0], 1ll * f[max(i - j - (j > 1), 0)] * dp[j - 1][1] % mod);
            if(i - j + 1 < b)
              getadd(dp[i][1], dp[j - 1][0]);
        }
    }
    getdec(ans, dp[n][0]);
    getdec(ans, dp[n][1]);
    for(int i = 1; i <= n; ++i)
      if(n - i + 1 < a)
        getdec(ans, 1ll * g[n - i + 1 - (i > 1)][1] * dp[i - 1][1] % mod);
    write(ans);
    return 0;
}