题解 P4690 [Ynoi2016] 镜中的昆虫

· · 题解

神仙题,orz lxl。

不带修

先考虑不带修怎么做,即 P1972,我们对每个点维护 pre_i 代表与它颜色相同且下标在它之前的,下标最大的点的下标。(好绕口)

那么我们对于一个区间询问,考虑对于每种颜色,只统计它在区间中出现的最左边一次,也就是查询有多少 l \leq i \leq r,满足 pre_i<l。那么我们可以作二维偏序来得到答案。

单点修

有了不带修的做法,单点修改的做法就很显然了。注意到进行一次单点修改对 pre 数组的影响是 O(1) 的。因此我们把一次单点修改变成一次删除和一次插入,这样就是在二维偏序的基础上加了时间维。因此我们可以作三维偏序。CDQ 分治即可。

区间修

先扔一个结论:进行所有区间修改后,pre 数组的单点修改次数是 O(n+m) 的。

凭啥?我们设 \{a_n\} 一个值域相同的极长连续段为一个「节点」。注意到如果一个节点整体赋上一个值,那么只有节点中第一个数的 pre 会被修改。

考虑 ODT 的过程,每次修改,我们先将两端的节点分裂,然后我们将中间的节点删掉,再换成同一个节点。

分裂和更换的过程,就是删除若干节点,再添加至多三个节点。对 pre 数组的修改次数和删除的节点个数同阶。而每个节点至多删除一次,我们添加的节点个数是 O(m) 的,初始的节点个数是 O(n) 的。因此,pre 数组的单点修改次数是 O(n+m) 的。

既然这样,我们就只需要找到 pre 数组被修改的位置和时间,然后按照单点修改的方式做就好了!

怎么找呢?其实就是上面复杂度分析那个过程。我们直接拿个 ODT 维护,然后每次修改就可以找出若干可能被修改的节点,然后对这些节点求修改后的 pre 即可。怎么求 pre?我们对每种颜色开个 ODT 即可。

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
int n, m;
map<int, int> mp;
int mcnt;
struct Node {
    int x, y, t, v, id;
} q[1500005];
bool cmpByT(const Node &a, const Node &b) {
    return a.t != b.t ? a.t < b.t : a.id < b.id;
}
bool cmpByX(const Node &a, const Node &b) {
    return a.x != b.x ? a.x < b.x : a.id < b.id;
}
int qcnt;
int pre[200005];
struct ODTNode {
    int l, r, v;
    bool operator<(const ODTNode &o) const {
        return l < o.l;
    }
};
struct ODT {
    typedef set<ODTNode>::iterator iter;
    set<ODTNode> tr, col[200005];
    iter Insert(int l, int r, int v) {
        col[v].insert({ l, r, v });
        return tr.insert({ l, r, v }).first;
    }
    void Delete(int l, int r, int v) {
        col[v].erase({ l, r, v });
        tr.erase({ l, r, v });
    }
    iter Split(int p) {
        iter it = tr.lower_bound({ p, 0, 0 });
        if (it != tr.end() && it->l == p) return it;
        it--;
        int l = it->l, r = it->r, v = it->v;
        Delete(l, r, v);
        Insert(l, p - 1, v);
        return Insert(p, r, v);
    }
    int Pre(int p) {
        iter it = --tr.upper_bound({ p, 0, 0 });
        if (it->l < p) return p - 1;
        else {
            iter co = col[it->v].lower_bound({ p, 0, 0 });
            if (co != col[it->v].begin()) return (--co)->r;
            return 0;
        }
    }
    void Assign(int l, int r, int v, int t) {
        iter itr = Split(r + 1), itl = Split(l);
        vector<int> ps;
        for (iter it = itl; it != itr; it++) {
            if (it != itl) ps.emplace_back(it->l);
            iter nxt = col[it->v].upper_bound(*it);
            if (nxt != col[it->v].end()) ps.emplace_back(nxt->l);
            col[it->v].erase(*it);
        }
        tr.erase(itl, itr);
        Insert(l, r, v);
        ps.emplace_back(l);
        iter nxt = col[v].upper_bound({ l, r, v });
        if (nxt != col[v].end()) ps.emplace_back(nxt->l);
        for (int i = 0; i < ps.size(); i++) {
            q[++qcnt] = { ps[i], pre[ps[i]], t, -1, 0 };
            pre[ps[i]] = Pre(ps[i]);
            q[++qcnt] = { ps[i], pre[ps[i]], t, 1, 0 };
        }
    }
} odt;
struct BIT {
    int f[100005];
    BIT() {
        memset(f, 0, sizeof(f));
    }
    int lowbit(int x) { return x & -x; }
    void Modify(int i, int x) {
        for (; i <= 100001; i += lowbit(i)) f[i] += x;
    }
    int Query(int i) {
        int res = 0;
        for (; i; i -= lowbit(i)) res += f[i];
        return res;
    }
} bit;
int a[100005], las[200005];
int ans[100005], acnt;
void CDQ(int l, int r) {
    if (l == r) return;
    int mid = l + r >> 1;
    CDQ(l, mid); CDQ(mid + 1, r);
    int i = l, j = mid + 1;
    while (j <= r) {
        while (i <= mid && q[i].x <= q[j].x) {
            if (!q[i].id) bit.Modify(q[i].y + 1, q[i].v);
            i++;
        }
        if (q[j].id) ans[q[j].id] += bit.Query(q[j].y + 1) * q[j].v;
        j++;
    }
    for (int k = l; k < i; k++) if (!q[k].id) bit.Modify(q[k].y + 1, -q[k].v);
    inplace_merge(q + l, q + mid + 1, q + r + 1, cmpByX);
}
int main() {
    scanf("%d%d", &n, &m);
    for (int i = 1; i <= n; i++) {
        scanf("%d", a + i);
        if (!mp[a[i]]) mp[a[i]] = ++mcnt;
        a[i] = mp[a[i]];
        pre[i] = las[a[i]];
        las[a[i]] = i;
        q[++qcnt] = { i, pre[i], 0, 1, 0 };
        odt.Insert(i, i, a[i]);
    }
    for (int i = 1; i <= m; i++) {
        int op;
        scanf("%d", &op);
        if (op == 1) {
            int l, r, v;
            scanf("%d%d%d", &l, &r, &v);
            if (!mp[v]) mp[v] = ++mcnt;
            v = mp[v];
            odt.Assign(l, r, v, i);
        }
        else {
            int l, r;
            scanf("%d%d", &l, &r);
            q[++qcnt] = { r, l - 1, i, 1, ++acnt };
            q[++qcnt] = { l - 1, l - 1, i, -1, acnt };
        }
    }
    sort(q + 1, q + qcnt + 1, cmpByT);
    CDQ(1, qcnt);
    for (int i = 1; i <= acnt; i++) {
        printf("%d\n", ans[i]);
    }
    return 0;
}

ODT 的实现参考了 Sol1 的题解。