如何使用 fhq Treap 通过 WBLT 的模板题

· · 题解

这里是点名被卡的 fhq Treap 做法,讲述了较多奇技淫巧和代码实现细节,如果想看到更详细的解析,请前往其他博客处查看。

下文中的 N=\mathcal{O}(n\log n) 指平衡树的节点个数。

首先这道题的四个操作看着就很能用平衡树来解决,有关 fhq Treap 的知识相信大家都会,在此略去。

笔者的代码用时为 739\text{ms},空间为 299.78\text{MB},下面将会讲述如何使用 fhq Treap 解决本题。

首先先写出一颗普通的 fhq Treap,然后在每次修改(即 splitmergepushdown)时都对于当前的版本进行一次备份,然后在备份上进行更改。

这个时候发现前两个操作很难处理,我们不妨定义一种“虚点”,其权值为 0 且不计入子树的大小(即如下代码中的 siz)中,也就是说,这棵 fhq Treap 是“半 Leafy 的”。

以下为 node 的定义:

struct op
{
    int ls, rs, siz, sum;
    bool rel, val, rev;
} p[N];

其中 rel 为真表示这个点是“实点”,反之则为“虚点”。

发现 relvalrev 都是 bool 变量,所以可以分别使用一个大小为 Nbitset 进行存储。

你可能会发现上面的定义没有常见的 key,因为在 fhq Treap 中,一种不需要比较键值的方式是用 \dfrac{siz_x}{siz_x+siz_y} 来加权比较 xy 的键值。

如果用了常见的 key 的方法的话,除了空间问题(带一个 \dfrac{5}{4} 的常数)外,还有可能会有由于复制操作带来的键值相同问题,而对于集合 S_1,S_2\max\limits_{x\in S_1}x>\max\limits_{x\in S_2}x 的概率为 \dfrac{|S_1|}{|S_1|+|S_2|},证明只需考虑 S_1\cup S_2 的最大值由哪个集合贡献而来即可。

现在直接按照别的题解一般实现无法做到优秀的复杂度,因为 fhq Treap 的虚点不能动态调整(也就是使用如下代码 Merge 虚点后再将 \mathcal{O}(\log k) 个“半成品”合并起来复杂度不正确)。

int Merge(int x, int y)
{
    int z = ++tot;
    p[z].ls = x, p[z].rs = y;
    rel.reset(z), val.reset(z), rev.reset(z);
    p[z].siz = p[x].siz + p[y].siz;
    p[z].sum = p[x].sum + p[y].sum;
    return z;
}

因为这可以通过复制再删除迭代的方式被搞炸,复杂度是错误的。

发现到对于给定的 k,我们求出最小的 k'=2^{x},然后将其复制 k' 编后截取前缀,这样只需使用到 1 个半成品。

但这样无法通过 lxl 精心设计的“不含操作 2 的测试点”,一个乱搞的方法是,在每次 1,2 操作后,随机选取一个位置,将序列分为两段,再将它们合并,最后的时间复杂度不会分析,但是能过。

参考代码:

#include <bits/stdc++.h>

using namespace std;

inline int read()
{
    int x;
    char c;
    while ((c = getchar()) < '0' || c > '9');
    x = c - '0';
    while ((c = getchar()) >= '0' && c <= '9') {
        x = (x << 3) + (x << 1) + (c ^ '0');
    }
    return x;
}

const int N = 2e7 + 5;

bitset<N> rel, val, rev;

struct op
{
    int ls, rs, siz, sum;
} p[N];

int rt, tot;
mt19937 rnd;

inline int newNode(int x)
{
    p[++tot].ls = 0;
    p[tot].rs = 0;
    p[tot].siz = 1;
    rel.set(tot);
    p[tot].sum = x;
    if (x) {
        val.set(tot);
    }
    return tot;
}

inline int backup(int x)
{
    if (!x) {
        return x;
    }
    p[++tot] = p[x];
    if (rel.test(x)) {
        rel.set(tot);
    }
    if (val.test(x)) {
        val.set(tot);
    }
    if (rev.test(x)) {
        rev.set(tot);
    }
    return tot;
}

inline void pushdown(int x)
{
    if (rev.test(x)) {
        p[x].ls = backup(p[x].ls);
        p[x].rs = backup(p[x].rs);
        rev.flip(p[x].ls);
        rev.flip(p[x].rs);
        swap(p[p[x].ls].ls, p[p[x].ls].rs);
        swap(p[p[x].rs].ls, p[p[x].rs].rs);
        rev.reset(x);
    }
    return;
}

int merge(int x, int y)
{
    if (!x || !y) {
        return x | y;
    }
    if ((int)((unsigned int)rnd() % (p[x].siz + p[y].siz)) < p[x].siz) {
        x = backup(x), pushdown(x);
        p[x].siz += p[y].siz;
        p[x].sum += p[y].sum;
        p[x].rs = merge(p[x].rs, y);
        return x;
    }
    y = backup(y), pushdown(y);
    p[y].siz += p[x].siz;
    p[y].sum += p[x].sum;
    p[y].ls = merge(x, p[y].ls);
    return y;
}

void split(int x, int k, int &a, int &b)
{
    if (!x) {
        a = b = 0;
        return;
    }
    x = backup(x);
    pushdown(x);
    if (p[p[x].ls].siz >= k) {
        b = x;
        split(p[x].ls, k, a, p[b].ls);
    } else {
        a = x;
        split(p[x].rs, k - p[p[x].ls].siz - rel.test(x), p[a].rs, b);
    }
    p[x].siz = p[p[x].ls].siz + p[p[x].rs].siz + rel.test(x);
    p[x].sum = p[p[x].ls].sum + p[p[x].rs].sum + val.test(x);
    return;
}

int Merge(int x, int y)
{
    int z = ++tot;
    p[z].ls = x, p[z].rs = y;
    rel.reset(z), val.reset(z), rev.reset(z);
    p[z].siz = p[x].siz + p[y].siz;
    p[z].sum = p[x].sum + p[y].sum;
    return z;
}

inline int power(int x, int y, int type)
{
    int len = p[x].siz * y, dt = 0, ty = y;
    y = 1;
    while (y < ty) {
        y <<= 1;
    }
    if (type == 1) {
        while (y) {
            if (y & 1) {
                dt = x;
            }
            x = Merge(x, x);
            y >>= 1;
        }
    } else {
        int c = x;
        x = backup(x);
        rev.flip(x);
        swap(p[x].ls, p[x].rs);
        int d = x;
        while (y) {
            if (y & 1) {
                dt = c;
            }
            int C = Merge(c, d), D = Merge(d, c);
            c = C, d = D;
            y >>= 1;
        }
    }
    if (y != ty) {
        int L, R;
        split(dt, len, L, R);
        dt = L;
    }
    return dt;
}

int getKth(int x, int y)
{
    if (val.test(x) && p[p[x].ls].sum + val.test(x) == y) {
        return p[p[x].ls].siz + 1;
    }
    pushdown(x);
    if (p[p[x].ls].sum >= y) {
        return getKth(p[x].ls, y);
    }
    return p[p[x].ls].siz + rel.test(x) + getKth(p[x].rs, y - p[p[x].ls].sum - val.test(x));
}

signed main()
{
    int n = read();
    char c;
    while ((c = getchar()) != '0' && c != '1');
    for (int i = 1; i <= n; ++i) {
        int x = c ^ '0';
        c = getchar();
        rt = merge(rt, newNode(x));
    }
    int m = read();
    while (m--) {
        int opt = read();
        if (opt <= 2) {
            int l = read(), r = read(), k = read(), x, y, z;
            split(rt, r, y, z);
            split(y, l - 1, x, y);
            y = power(y, k, opt);
            rt = merge(merge(x, y), z);
            if (p[rt].siz != 1) {
                int x, y;
                int pos = rnd() % (p[rt].siz - 1) + 1;
                split(rt, pos, x, y);
                rt = merge(x, y);
            }
        } else if (opt == 3) {
            int l = read(), r = read(), x, y, z;
            split(rt, r, y, z);
            split(y, l - 1, x, y);
            rt = merge(x, z);
        } else if (opt == 4) {
            int k = read();
            if (k > p[rt].sum) {
                puts("-1");
                continue;
            }
            printf("%d\n", getKth(rt, k));
        }
    }
    return 0;
}

有没有人能证明或者能卡掉啊 /kel