P10008 [集训队互测 2022] Range Minimum Element

· · 题解

或许更好的阅读体验。

题意:

有一个长度为 n,值域为 [1,c] 的正整数序列 a。给定 m 个区间 [l_i, r_i],设长度为 m 的序列 b 满足 \forall i \in [1,m], b_i = \min\limits_{j=l_i}^{r_i}\{a_j\}

求出 a 在范围内任意取的情况下共能得到多少种不同的 b。答案对 998244353 取模。

其中 1 \leq n \leq 1001 \leq m \leq \dfrac{n(n + 1)}{2}1 \leq c < 998244353\forall i \in [1,m], 1 \leq l_i \leq r_i \leq n

思路:

这种题,直接做是困难的,就是考虑构造一组 b \to a 的单射,然后统计 a 即可,显然这个 a 是加了限制条件的。

对于一组固定的 b,其对应的 a 初始为空,然后考虑按照值 v 从大到小考虑每个区间 [l, r],将这个区间内非空的位置全部赋值为 v,如果没有非空的位置,则无解;最后为空的位置赋值为 1;容易发现,这是一组单射。

于是一个 a 是合法的,可以通过子问题划分:

于是一个 a 是否合法可以找里面第一个最小值位置然后划分到左右是否合法。

那么可以想到 dp 的状态,即 dp_{i, l, r} 表示 a 的区间 [l, r][i, c] 范围内的数的合法序列数,然后设一个辅助数组 f_{l, r} 表示 [l, r] 是否被 [l, r] 内的区间完美覆盖,那么可以得到转移:

dp_{i, l, r} \gets dp_{i, l,r } + f_{l, r} \cdot dp_{i + 1, l, r} dp_{i, l, r} \gets dp_{i, l, r} + \sum_{k = l}^r f_{l, k - 1} \cdot dp_{i + 1, l, k - 1} \cdot dp_{i, k + 1, r}

这样做是 O(cn^3) 的,考虑优化;发现 c 很大,又不容易从状态里去掉,于是可以猜测 f_{i, l, r} 是关于 i 的一个多项式。

你手摸一下式子 f_{i, l, l} = c - i + 1 是一次的,于是可以归纳得到 f_{i, l, r} 是关于 i 的不高于 r - l + 1 次函数。

于是只需要算出最大的 n + 1if_{i, 1, n},然后拉格朗日插值即可。

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

完整代码:

#include<bits/stdc++.h>
#define ls(k) k << 1
#define rs(k) k << 1 | 1
#define lowbit(x) x & (-x)
#define fi first
#define se second
#define popcnt(x) __builtin_popcount(x)
#define open(s1, s2) freopen(s1, "r", stdin), freopen(s2, "w", stdout);
using namespace std;
typedef __int128 __;
typedef long double lb;
typedef double db;
typedef unsigned long long ull;
typedef long long ll;
bool Begin;
const int N = 105, M = 1e4 + 10, mod = 998244353;
inline ll read() {
    ll x = 0, dp = 1;
    char c = getchar();
    while (c < '0' || c > '9') {
        if (c == '-')
            dp = -1;
        c = getchar();
    }
    while (c >= '0' && c <= '9') {
        x = (x << 1) + (x << 3) + (c ^ 48);
        c = getchar();
    }
    return x * dp;
}
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 void getdec(int &x, int y){
    x = (x < y) ? (x - y + mod) : (x - y);
}
int n, m, c;
int d[N], x[N], y[N], L[M], R[M];
int dp[N][N][N];
bool f[N][N];
inline int qpow(int a, int b){
    int ans = 1;
    while(b){
        if(b & 1)
            ans = 1ll * ans * a % mod;
        a = 1ll * a * a % mod;
        b >>= 1; 
    }
    return ans;
}
inline int getf(int n, int k){
    int sum = 0;
    for(int i = 1; i <= n; ++i){
        int a = 1, b = 1;
        for(int j = 1; j <= n; ++j){
            if(i == j)
                continue;
            a = 1ll * (k - x[j] + mod) % mod * a % mod;
            b = 1ll * (x[i] - x[j] + mod) % mod * b % mod;
        }
        sum = (sum + 1ll * y[i] * a % mod * qpow(b, mod - 2) % mod) % mod; 
    }
    return sum;
}
int main() {
    n = read(), m = read(), c = read();
    for(int i = 1; i <= m; ++i)
      L[i] = read(), R[i] = read();
    for(int l = 1; l <= n; ++l){
        for(int r = l; r <= n; ++r){
            for(int i = l - 1; i <= r + 1; ++i)
              d[i] = 0;
            for(int i = 1; i <= m; ++i)
              if(l <= L[i] && R[i] <= r)
                ++d[L[i]], --d[R[i] + 1];
            f[l][r] = 1;
            for(int i = l; i <= r; ++i){
                d[i] += d[i - 1];
                if(!d[i]){
                    f[l][r] = 0;
                    break;
                }
            }
//          cerr << f[l][r];
        }
    }
//  cerr << '\n';
//  cerr << f[1][2] << ' ' << f[1][1] << '\n';
    for(int i = 1; i <= n + 1; ++i){
//      cerr << i << '\n';
        x[i] = c - i + 1;
        for(int len = 1; len <= n; ++len){
            for(int l = 1; l + len - 1 <= n; ++l){
                int r = l + len - 1;
                if(f[l][r])
                  getadd(dp[i][l][r], dp[i - 1][l][r]);
                if(l == r)
                  getadd(dp[i][l][r], 1);
                else
                  getadd(dp[i][l][r], dp[i][l + 1][r]);
                if(f[l][r - 1])
                  getadd(dp[i][l][r], dp[i - 1][l][r - 1]);
                for(int k = l + 1; k < r; ++k)
                  if(f[l][k - 1])
                    getadd(dp[i][l][r], 1ll * dp[i - 1][l][k - 1] * dp[i][k + 1][r] % mod);
//              cerr << "   " << l << ' ' << r << ' ' << dp[i][l][r] << '\n';
            }
//          cerr << '\n';
        }
//      cerr << '\n' << '\n';
        y[i] = dp[i][1][n];
//      cerr << y[i] << '\n';
    }
    if(c <= n + 1){
        write(y[c]);
        putchar('\n');
        return 0;
    }
    write(getf(n + 1, 1));
    return 0;
}