P8996 [CEOI 2022] Abracadabra

· · 题解

很好的题,场切了。

思路:

考虑一次归并本质是在做什么,先划分为 L, R 两段,对于里面的任何一个数 a_i,如果 a_i 被插入进去了,那么在 a_i 后面 < a_i 的整个段也会一起被插入进去。

于是你发现了一个关键性质,将 L, R 中划分为若干段,每段由一个 a_i 以及它后面连续小于 a_i 的极长段构成,那么一次归并本质上是按照这些段的段首排序,然后拼接起来;但是排序后可能会有段跨过了中间,需要先从中间劈开,然后对右边劈开的段再划分一下。

接下来考虑怎么维护上面的东西,因为是按照段首值排序,考虑权值线段树维护 len_i 表示以 i 开头的段的长度,那么排序一次是首先先找到一个 i 使得它这个段跨过了中间 \frac{n}{2},设为 [l, r];接下来考虑怎么快速划分 [\frac{n}{2} + 1, r]:

询问显然可以离线放到每个时刻上去,对于一个 t 时刻的询问 x,可以找到最小的 u 使得 \sum_{i \le u} len_i \ge x,那么 xu 开头的段中,于是就可以确认 x 位置是哪个元素了。

上面找跨过 \frac{n}{2} 的段与找到 x 所在段都可以在线段树上二分简单实现,时间复杂度为 O(n \log n)

完整代码:

#include<bits/stdc++.h>
#define fi first

#define se second

using namespace std;
typedef long long ll;
const int N = 2e5 + 10, M = 1e6 + 10;
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');
}
struct Node{
    int l, r;
    int sum;
}X[N << 2];
int n, q, t, x, top;
int p[N], pp[N], nxt[N], stk[N], len[N], ans[M];
vector<pair<int, int>> Q[N];
inline void pushup(int k){
    X[k].sum = X[k << 1].sum + X[k << 1 | 1].sum;
}
inline void build(int k, int l, int r){
    X[k].l = l, X[k].r = r;
    if(l == r)  
      return ;
    int mid = (l + r) >> 1;
    build(k << 1, l, mid);
    build(k << 1 | 1, mid + 1, r);
}
inline void update(int k, int i, int v){
    if(X[k].l == i && i == X[k].r){
        X[k].sum = v;
        return ;
    }
    int mid = (X[k].l + X[k].r) >> 1;
    if(i <= mid)
      update(k << 1, i, v);
    else
      update(k << 1 | 1, i, v);
    pushup(k);
}
inline int getk(int k, int v, int &sum){
    if(X[k].l == X[k].r){
        sum += X[k].sum;
        return X[k].l;
    }
    int mid = (X[k].l + X[k].r) >> 1;
    if(sum + X[k << 1].sum >= v)
      return getk(k << 1, v, sum);
    else{
        sum += X[k << 1].sum;
        return getk(k << 1 | 1, v, sum);
    }
}
int main(){
    // freopen("magic.in", "r", stdin);
    // freopen("magic.out", "w", stdout);
    n = read(), q = read();
    for(int i = 1; i <= n; ++i){
        p[i] = read();
        pp[p[i]] = i;
        while(top && p[i] > p[stk[top]]){
            nxt[stk[top]] = i;
            --top;
        }
        stk[++top] = i;
    }
    while(top)
      nxt[stk[top--]] = n + 1;
    build(1, 1, n);
    for(int i = 1; i <= n;){
        if(i <= (n >> 1) && nxt[i] - 1 > (n >> 1)){
            len[p[i]] = (n >> 1) - i + 1;
//          cerr << p[i] << ' ' << len[p[i]] << '\n';
            update(1, p[i], len[p[i]]);
            i = (n >> 1) + 1;
        } 
        else{
            len[p[i]] = nxt[i] - i;
//          cerr << p[i] << ' ' << len[p[i]] << '\n';
            update(1, p[i], len[p[i]]);
            i = nxt[i];         
        }
    }
    for(int i = 1; i <= q; ++i){
        t = read(), x = read();
        if(!t)
          ans[i] = p[x];
        else
          Q[min(t, n)].push_back({x, i});
    }
    bool flag = 0;
    for(int tim = 1; tim <= n; ++tim){
        for(auto t : Q[tim]){
            int pos = t.fi, id = t.se;
            int end = 0;
            int u = getk(1, pos, end);
            int start = end - len[u] + 1;
            ans[id] = p[pp[u] + pos - start];
        }
        if(!flag){
            int end = 0;
            int u = getk(1, (n >> 1) + 1, end);
//          cerr << u << ' ' << end << '\n';
            int start = end - len[u] + 1;
            if(start > (n >> 1)){
                flag = 1;
                continue;
            }
            int now = pp[u] + (n >> 1) + 1 - start;
//          cerr << start << ' ' << now << '\n'; 
            for(int i = now; i <= pp[u] + len[u] - 1;){
                len[p[i]] = min(nxt[i], pp[u] + len[u])  - i;
//              cerr << "new: " << p[i] << ' ' << len[p[i]] << '\n';
                update(1, p[i], len[p[i]]);
                i = nxt[i];
            }
            len[u] = (n >> 1) - start + 1;
            update(1, u, len[u]);
        }
    }
    for(int i = 1; i <= q; ++i){
        write(ans[i]);
        putchar('\n');
    }
    return 0;
}