Splay 85pts TLE on #11~13 求助
查看原帖
Splay 85pts TLE on #11~13 求助
182234
ryanright楼主2023/7/20 16:50

rt,下载数据后在本地测试了一下,代码在运行时跑了一会跑着跑着就不动了,哪位大佬能帮忙调一下?

#include <cstdio>
#include <utility>
#define int long long
using namespace std;
int n, m, cnt;
struct node {
    int son[2], fa, sub, l, r;
} tr[2000005];
struct Splay {
    int root;
    inline void pushup(int x) {
        if (!x)
            return;
        tr[x].sub = tr[x].r - tr[x].l + 1;
        if (tr[x].son[0])
            tr[x].sub += tr[tr[x].son[0]].sub;
        if (tr[x].son[1])
            tr[x].sub += tr[tr[x].son[1]].sub;
    }
    inline bool check(int x) {
        return tr[tr[x].fa].son[1] == x;
    }
    inline void rorate(int x) {
        int y = tr[x].fa, z = tr[y].fa;
        bool w = check(x);
        tr[y].son[w] = tr[x].son[w ^ 1];
        tr[tr[y].son[w]].fa = y;
        tr[y].fa = x;
        tr[x].son[w ^ 1] = y;
        tr[x].fa = z;
        if (z)
            tr[z].son[tr[z].son[1] == y] = x;
        pushup(y);
        pushup(x);
    }
    void splay(int x, int goal) {
        for (int f; (f = tr[x].fa) != goal; rorate(x))
            if (tr[f].fa != goal)
                rorate(check(x) == check(f) ? f : x);
        if (!goal)
            root = x;
    }
    inline int new_node(int l, int r) {
        cnt++;
        tr[cnt].sub = (tr[cnt].r = r) - (tr[cnt].l = l) + 1;
        return cnt;
    }
    inline void init(int l, int r) {
        root = new_node(l, r);
    }
    int next_node(int x) {
        splay(x, 0);
        int pos = tr[x].son[1];
        while (tr[pos].son[0])
            pos = tr[pos].son[0];
        return pos;
    }
    int last_node(int x) {
        splay(x, 0);
        int pos = tr[x].son[0];
        while (tr[pos].son[1])
            pos = tr[pos].son[1];
        return pos;
    }
    int split(int x, int k) { // Split x into [l, k] and (k, r].
        if (k >= tr[x].r || k < tr[x].l)
            return x;
        int y = new_node(k + 1, tr[x].r);
        tr[x].r = k;
        if (tr[x].son[1]) {
            int pos = tr[x].son[1];
            while (tr[pos].son[0])
                pos = tr[pos].son[0];
            tr[tr[pos].son[0] = y].fa = pos;
            while (pos != x) {
                pushup(pos);
                pos = tr[pos].fa;
            }
        } else
            tr[tr[x].son[1] = y].fa = x;
        splay(y, 0);
        return y;
    }
    pair<int, int> find(int x) { // Find the kth number and the node includes it in the splay.
        int pos = root;
        while (true) {
            if (x <= tr[tr[pos].son[0]].sub)
                pos = tr[pos].son[0];
            else {
                x -= tr[tr[pos].son[0]].sub;
                if (x <= tr[pos].r - tr[pos].l + 1)
                    return make_pair(pos, tr[pos].l + x - 1);
                x -= tr[pos].r - tr[pos].l + 1;
                pos = tr[pos].son[1];
            }
        }
    }
    void pop(int x) { // Delete the kth node in the splay.
        int nxt = next_node(x), lst = last_node(x);
        if (!nxt) {
            tr[tr[x].son[0]].fa = 0;
            root = tr[x].son[0];
        } else if (!lst) {
            tr[tr[x].son[1]].fa = 0;
            root = tr[x].son[1];
        } else {
            splay(nxt, 0);
            splay(lst, root);
            tr[lst].son[1] = 0;
            pushup(lst);
            pushup(nxt);
        }
        tr[x].fa = tr[x].son[0] = tr[x].son[1] = 0;
        tr[x].sub = tr[x].r - tr[x].l + 1;
    }
    void push_back(int x) { // Push node x to the back of the splay.
        int pos = root;
        while (tr[pos].son[1])
            pos = tr[pos].son[1];
        tr[tr[pos].son[1] = x].fa = pos;
        while (pos) {
            pushup(pos);
            pos = tr[pos].fa;
        }
        splay(x, 0);
    }
} s[2000005];
signed main() {
    int n, m, q;
    scanf("%lld%lld%lld", &n, &m, &q);
    for (int i = 1; i <= n; i++)
        s[i].init((i - 1) * m + 1, i * m - 1);
    s[0].init(m, m);
    for (int i = 2; i <= n; i++)
        s[0].push_back(s[0].new_node(i * m, i * m));
    while (q--) {
        int x, y;
        scanf("%lld%lld", &x, &y);
        if (y == m) {
            auto val = s[0].find(x);
            printf("%lld\n", val.second);
            s[0].pop(val.first);
            s[0].push_back(val.first);
            continue;
        }
        auto val = s[x].find(y);
        printf("%lld\n", val.second);
        s[x].split(s[x].split(val.first, val.second - 1), val.second);
        int pos = s[x].find(y).first;
        s[x].pop(pos);
        s[0].push_back(pos);
        val = s[0].find(x);
        s[0].split(s[0].split(val.first, val.second - 1), val.second);
        pos = s[0].find(x).first;
        s[0].pop(pos);
        s[x].push_back(pos);
    }
    return 0;
}
2023/7/20 16:50
加载中...