P9129 [USACO23FEB] Piling Papers G

· · 题解

或许更好的阅读体验。

题意:

N 张写着数字的纸片排成一排,给定 A,B

希望你回答 Q 个查询。每次查询,将左到右遍历纸片 lr。有一个最初为空的纸片堆。遍历到张纸片,它们可以选择将其添加到堆的顶部、底部,或者不添加。最后,它们将从顶部到底部读取堆中的纸片,形成一个整数。在奶牛们在此过程中做选择的所有 3^{r_i - l_i + 1} 种方式中,计算出结果在 [A,B] 范围内的方式数量,并输出这个数量对 10^9 + 7 取模的结果。

其中 N \leq 300Q \leq 5 \times 10^4

思路:

先考虑全局 l = 1, r = n 怎么做。

首先差分一下 [A, B] = [1, A] - [1, B - 1],设我们现在要求 \le x 的答案。

然后考虑数位 dp,发现直接做不太行,因为一个数既可以插在前面,还可以插在后面,如何只记录前缀的话,那么插在前面是不容易判断大小的。

所以考虑区间形式的状态,定义 dp_{i, l, r, 0/1/2} 表示考虑前 i 个数,其中选择了 r - l + 1 个数,对应 x[l, r] 数位,与其的相对大小关系是小于,等于,大于的方案数。

x 的位数是 m,那么答案容易发现是:

\sum_{i = 1}^{m} dp_{n, i, m, 0} + dp_{n, i, m, 1} + [i > 1] dp_{n, i, m, 2}

即考虑最终形成的位数,当位数等于 m 时,不能大于 x,位数小于 m 时,可以任意。

转移时考虑 a_j 放在左右两边的情况即可。

这样时间复杂度可以做到 O(q n \log^2 w),无法通过,因为 q 比较大,考虑提前预处理出每个区间的答案,即可 O(1) 查询,时间复杂度为 O(n^2 \log^2 w + q)

完整代码:

#include<bits/stdc++.h>
#define lowbit(x) x & (-x)
#define pi pair<ll, ll>
#define ls(k) k << 1
#define rs(k) k << 1 | 1
#define fi first
#define se second
using namespace std;
typedef __int128 __;
typedef long double lb;
typedef double db;
typedef unsigned long long ull;
typedef long long ll;
const int N = 305, 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');
}
ll A, B;
int n, m, q, l, r;
int a[N], h[N];
int dp[N][N][3];
int ans[2][N][N];
inline void init(ll x){
    m = 0;
    while(x){
        h[++m] = x % 10;
        x /= 10;
    }
    reverse(h + 1, h + m + 1);
}
inline int get(int x, int y){
    if(x < y)
      return 0;
    else if(x == y)
      return 1;
    else
      return 2;
}
inline int add(int x, int y){
    return (x + y >= mod) ? (x + y - mod) : (x + y); 
}
inline void getadd(int &x, int y){
    x = add(x, y);
}
inline void solve(ll x, int id){
    init(x);
    for(int i = 1; i <= n; ++i){
        memset(dp, 0, sizeof(dp));
        for(int j = i; j <= n; ++j){
            for(int l = 1; l <= m; ++l){
                for(int r = m; r > l; --r){
                    if(a[j] > h[l]){
                        getadd(dp[l][r][2], dp[l + 1][r][0]);
                        getadd(dp[l][r][2], dp[l + 1][r][1]);
                        getadd(dp[l][r][2], dp[l + 1][r][2]);
                    }
                    else if(a[j] == h[l]){
                        getadd(dp[l][r][0], dp[l + 1][r][0]);
                        getadd(dp[l][r][1], dp[l + 1][r][1]);
                        getadd(dp[l][r][2], dp[l + 1][r][2]); 
                    }
                    else{
                        getadd(dp[l][r][0], dp[l + 1][r][0]);
                        getadd(dp[l][r][0], dp[l + 1][r][1]);
                        getadd(dp[l][r][0], dp[l + 1][r][2]);
                    }
                    getadd(dp[l][r][0], dp[l][r - 1][0]);
                    getadd(dp[l][r][2], dp[l][r - 1][2]); 
                    getadd(dp[l][r][get(a[j], h[r])], dp[l][r - 1][1]);
                }
            }
            for(int x = 1; x <= m; ++x)
              getadd(dp[x][x][get(a[j], h[x])], 2);
            for(int x = 1; x <= m; ++x){
                getadd(ans[id][i][j], dp[x][m][0]);
                getadd(ans[id][i][j], dp[x][m][1]);
                if(x > 1)
                  getadd(ans[id][i][j], dp[x][m][2]);
            }
        }
    }
}
int main(){
    n = read(), A = read() - 1, B = read();
    for(int i = 1; i <= n; ++i)
      a[i] = read();
    solve(A, 0);
    solve(B, 1);
    q = read();
    while(q--){
        l = read(), r = read();
        write((ans[1][l][r] - ans[0][l][r] + mod) % mod);
        putchar('\n');
    }
    return 0;
}