题解:P17154 [ICPC 2017 Xi'an R] Acedia

· · 题解

有点意思哈。

考虑用莫队维护当前区间内有哪些数字出现过。

维护一个数组 cnt[x],表示数字 x 在当前区间内出现了几次。

当我们往区间里加一个数字 v 时,有三种情况:

删除同理,直接反过来即可。

然后考虑怎么知道一个段的长度,其实很简单,我们只需要统计长度不超过 10 的段。所以当加入数字 v 时,只需要往左看最多 10 个连续的数字,往右看最多 10 个连续的数字,就知道左右段的长度了。

最后讲讲代码逻辑:

由于值域很大,先离散化一下,同时预处理每个数的“邻居”,用 pre[i] 表示 i-1 的下标(如果没有就是 -1),用 nxt[i] 表示 i+1 的下标(如果没有就是 -1)。

然后是添加数字和删除数字的部分,直接看代码吧:

void add_val(int id) {
    int l = get_lft(id), r = get_rgt(id); // 左边和右边连续段长度
    // 减去被合并的段
    if (1 <= l && l <= 10) ans[l]--;
    if (1 <= r && r <= 10) ans[r]--;
    // 加上合并后的新段
    int tmp = l + r + 1;
    if (1 <= tmp && tmp <= 10) ans[tmp]++;
}
void del_val(int id) {
    int l = get_lft(id), r = get_rgt(id); // 左边和右边连续段长度
    // 加上合并后的新段
    if (1 <= l && l <= 10) ans[l]++;
    if (1 <= r && r <= 10) ans[r]++;
    // 减去原来的整段
    int tmp = l + r + 1;
    if (1 <= tmp && tmp <= 10) ans[tmp]--;
}

再看如何获取连续段长度,这边以左边的连续段为例:

int get_lft(int id) {
    int len = 0, cur = pre[id]; // 当前长度 当前查到的数字的下标 
    while (cur != -1 && cnt[cur] > 0 && len < 10) len++, cur = pre[cur]; // cur 没有到边界并且出现在了当前区间内并且长度实在我们关心的范围内(不到10)长度增加,下标更新 
    if (len == 10 && cur != -1 && cnt[cur] > 0) return 11; // 如果长度到了上限但还能蒸,直接标个11结束 
    return len; // 否则返回长度 
}

然后就很板了,看最终代码:

#include <bits/stdc++.h>
#define Write ios::sync_with_stdio(0);
#define by cin.tie(0);
#define Na1L0n9 cout.tie(0);
using namespace std;
typedef long long ll;
typedef unsigned long long ull;
const int N = 1e6 + 10;
const int MOD = 998244353;
int T, n, m, vid, a[N], val[N], id[N], pre[N], nxt[N], cnt[N], ans[20];
struct Question {
    int l, r, id, bk;
    bool operator < (const Question& other) const {
        if (bk != other.bk) return bk < other.bk;
        return (bk & 1 ? r < other.r : r > other.r);
    }
} q[N];
char out[N][20];
int get_lft(int id) {
    int len = 0, cur = pre[id];
    while (cur != -1 && cnt[cur] > 0 && len < 10) len++, cur = pre[cur];
    if (len == 10 && cur != -1 && cnt[cur] > 0) return 11;
    return len;
}
int get_rgt(int id) {
    int len = 0, cur = nxt[id];
    while (cur != -1 && cnt[cur] > 0 && len < 10) len++, cur = nxt[cur];
    if (len == 10 && cur != -1 && cnt[cur] > 0) return 11;
    return len;
}
void add_val(int id) {
    int l = get_lft(id), r = get_rgt(id);
    if (1 <= l && l <= 10) ans[l]--;
    if (1 <= r && r <= 10) ans[r]--;
    int tmp = l + r + 1;
    if (1 <= tmp && tmp <= 10) ans[tmp]++;
}
void del_val(int id) {
    int l = get_lft(id), r = get_rgt(id);
    if (1 <= l && l <= 10) ans[l]++;
    if (1 <= r && r <= 10) ans[r]++;
    int tmp = l + r + 1;
    if (1 <= tmp && tmp <= 10) ans[tmp]--;
}
void add_pos(int p) {
    cnt[id[p]]++;
    if (cnt[id[p]] == 1) add_val(id[p]);
}
void del_pos(int p) {
    cnt[id[p]]--;
    if (!cnt[id[p]]) del_val(id[p]);
}

int main() {
    Write by Na1L0n9
    cin >> T;
    while (T--) {
        cin >> n >> m;
        for (int i = 1; i <= n; i++) {
            cin >> a[i];
            val[i] = a[i];
        }
        sort(val + 1, val + n + 1);
        vid = unique(val + 1, val + n + 1) - (val + 1);
        for (int i = 1; i <= n; i++) id[i] = lower_bound(val + 1, val + vid + 1, a[i]) - (val + 1);
        for (int i = 1; i <= vid; i++) {
            pre[i - 1] = -1, nxt[i - 1] = -1;
            if (i > 1 && val[i - 1] == val[i] - 1) pre[i - 1] = i - 2;
            if (i < vid && val[i + 1] == val[i] + 1) nxt[i - 1] = i;
        }
        int sz = max(1, (int)(n / sqrt(m + 1)) + 1);
        for (int i = 0; i < m; i++) cin >> q[i].l >> q[i].r, q[i].id = i, q[i].bk = q[i].l / sz;
        sort(q, q + m);
        fill(cnt, cnt + vid, 0); fill(ans, ans + 11, 0);
        int cl = 1, cr = 0;
        for (int i = 0; i < m; ++i) {
            int L = q[i].l, R = q[i].r;
            while (cl > L) add_pos(--cl);
            while (cr < R) add_pos(++cr);
            while (cl < L) del_pos(cl++);
            while (cr > R) del_pos(cr--);
            for (int k = 1; k <= 10; k++) out[q[i].id][k - 1] = '0' + (ans[k] % 10);
            out[q[i].id][10] = '\n', out[q[i].id][11] = '\0';
        }
        //for (int i = 0; i < m; i++) cout << out[i]; 千万~不要~这么~写啊~~~
        for (int i = 0; i < m; i++) fwrite(out[i], 1, 11, stdout);
    }

    return 0;
}