蒟蒻求助线段树套线段树重复利用节点
查看原帖
蒟蒻求助线段树套线段树重复利用节点
357440
NullNone楼主2023/7/19 17:22

动态开点线段树用了400M内存

所以想把没用的节点回收,用了一个叫做 mem 的栈来储存没用的节点的编号,然后挂了...

第66行 // mem.push(cur);就是回收节点...

#include <iostream>
#include <vector>
#include <stack>
#include <map>
#include <algorithm>
using namespace std;
const int MAXN = 5e4 + 5;
const int INF = 2147483647;
int n, m, a[MAXN];
struct Segment_Tree{
    const int llim = 0, rlim = 1e5;
    vector<int>lson = {0, 0};
    vector<int>rson = {0, 0};
    vector<int>siz = {0, 0};
    stack<int>mem;
    inline int node(){
        if(!mem.empty()){
            int cur = mem.top();
            mem.pop();
            lson[cur] = rson[cur] = siz[cur] = 0;
            return cur;
        }
        int cur = lson.size();
        lson.push_back(0);
        rson.push_back(0);
        siz.push_back(0);
        return cur;
    }
    // inline void update(int cur){
    //     siz[cur] = 0;
    //     if(lson[cur])
    //         siz[cur] += siz[lson[cur]];
    //     if(rson[cur])
    //         siz[cur] += siz[rson[cur]];
    // }
    int insert(int cur, int lt, int rt, int pos){
        if(!cur)
            cur = node();
        ++siz[cur];
        if(lt == rt){
            return cur;
        }
        int mid = (lt + rt) >> 1;
        if(pos <= mid)
            lson[cur] = insert(lson[cur], lt, mid, pos);
        else
            rson[cur] = insert(rson[cur], mid + 1, rt, pos);
        // update(cur);
        return cur;
    }
    inline void insert(int pos){insert(1, llim, rlim, pos);}
    int del(int cur, int lt, int rt, int pos){
        if(!cur)
            return 0;
        --siz[cur];
        if(lt == rt){
        }else{
            int mid = (lt + rt) >> 1;
            if(pos <= mid)
                lson[cur] = del(lson[cur], lt, mid, pos);
            else
                rson[cur] = del(rson[cur], mid + 1, rt, pos);
            // update(cur);
        }
        if(!siz[cur]){
            // mem.push(cur);
            return 0;
        }
        return cur;
    }
    inline void del(int pos){del(1, llim, rlim, pos);}
    int query(int cur, int lt, int rt, int st, int en){
        if(!cur)
            return 0;
        if(st <= lt && rt <= en)
            return siz[cur];
        int mid = (lt + rt) >> 1;
        int res = 0;
        if(st <= mid)
            res = query(lson[cur], lt, mid, st, en);
        if(en > mid)
            res += query(rson[cur], mid + 1, rt, st, en);
        return res;
    }
    inline int query(int st, int en){return query(1, llim, rlim, st, en);}
    int prev(int cur, int lt, int rt, int pos){
        if(!cur)
            return -INF;
        if(lt == rt)
            return lt;
        int mid = (lt + rt) >> 1;
        if(pos <= mid + 1)
            return prev(lson[cur], lt, mid, pos);
        int res = prev(rson[cur], mid + 1, rt, pos);
        if(res == -INF)
            return prev(lson[cur], lt, mid, pos);
        return res;
    }
    inline int prev(int pos){return prev(1, llim, rlim, pos);}
    int next(int cur, int lt, int rt, int pos){
        if(!cur)
            return INF;
        if(lt == rt)
            return lt;
        int mid = (lt + rt) >> 1;
        if(pos >= mid)
            return next(rson[cur], mid + 1, rt, pos);
        int res = next(lson[cur], lt, mid, pos);
        if(res == INF)
            return next(rson[cur], mid + 1, rt, pos);
        return res;
    }
    inline int next(int pos){return next(1, llim, rlim, pos);}
};
struct T2_Segment_Tree{
    Segment_Tree ts[MAXN << 2];
    void insert(int cur, int lt, int rt, int pos, int val){
        ts[cur].insert(val);
        if(lt == rt)
            return;
        int mid = (lt + rt) >> 1;
        if(pos <= mid)
            insert(cur << 1, lt, mid, pos, val);
        else
            insert(cur << 1 | 1, mid + 1, rt, pos, val);
    }
    inline void insert(int pos, int val){insert(1, 1, n, pos, val);}
    void del(int cur, int lt, int rt, int pos, int val){
        ts[cur].del(val);
        if(lt == rt)
            return;
        int mid = (lt + rt) >> 1;
        if(pos <= mid)
            del(cur << 1, lt, mid, pos, val);
        else
            del(cur << 1 | 1, mid + 1, rt, pos, val);
    }
    inline void del(int pos, int val){del(1, 1, n, pos, val);}
    int prev(int cur, int lt, int rt, int st, int en, int val){
        if(st <= lt && rt <= en)
            return ts[cur].prev(val);
        int mid = (lt + rt) >> 1;
        int res = -INF;
        if(st <= mid)
            res = prev(cur << 1, lt, mid, st, en, val);
        if(en > mid)
            res = max(res, prev(cur << 1 | 1, mid + 1, rt, st, en, val));
        return res;
    }
    inline int prev(int st, int en, int val){return prev(1, 1, n, st, en, val);}
    int next(int cur, int lt, int rt, int st, int en, int val){
        if(st <= lt && rt <= en)
            return ts[cur].next(val);
        int mid = (lt + rt) >> 1;
        int res = INF;
        if(st <= mid)
            res = next(cur << 1, lt, mid, st, en, val);
        if(en > mid)
            res = min(res, next(cur << 1 | 1, mid + 1, rt, st, en, val));
        return res;
    }
    inline int next(int st, int en, int val){return next(1, 1, n, st, en, val);}
    int query(int cur, int lt, int rt, int st, int en, int val){
        if(st <= lt && rt <= en)
            return ts[cur].query(0, val);
        int mid = (lt + rt) >> 1;
        int res = 0;
        if(st <= mid)
            res = query(cur << 1, lt, mid, st, en, val);
        if(en > mid)
            res += query(cur << 1 | 1, mid + 1, rt, st, en, val);
        return res;
    }
    inline int query(int st, int en, int val){return query(1, 1, n, st, en, val);}
    inline int rank(int st, int en, int k){
        int lt = 0, rt = 100000, mid;
        while(lt < rt - 1){
            mid = (lt + rt) >> 1;
            if(query(st, en, mid - 1) < k)
                lt = mid;
            else
                rt = mid;
        }
        return lt;
    }
}tnt;
int opt[MAXN], l[MAXN], r[MAXN], ps[MAXN], rk[MAXN], x[MAXN], res;
map<int, int>nval;
vector<int>rval;
int main(int argc, char const *argv[])
{
    ios::sync_with_stdio(false);
    cin >> n >> m;
    for(int i = 1; i <= n; ++i){
        cin >> a[i];
        rval.push_back(a[i]);
    }
    for(int i = 0; i < m; ++i){
        cin >> opt[i];
        switch(opt[i]){
            case 1: cin >> l[i] >> r[i] >> x[i]; break;
            case 2: cin >> l[i] >> r[i] >> rk[i]; break;
            case 3: cin >> ps[i] >> x[i]; break;
            case 4: cin >> l[i] >> r[i] >> x[i]; break;
            case 5: cin >> l[i] >> r[i] >> x[i]; break;
        }
        if(opt[i] != 2)
            rval.push_back(x[i]);
    }
    sort(rval.begin(), rval.end());
    rval.erase(unique(rval.begin(), rval.end()), rval.end());
    for(int i = 0; i < rval.size(); ++i)
        nval[rval[i]] = i;
    for(int i = 1; i <= n; ++i){
        a[i] = nval[a[i]];
        tnt.insert(i, a[i]);
        // cerr << a[i] << ' ';
    }
    // cerr << endl;
    for(int i = 0; i < m; ++i){
        if(opt[i] != 2){
            x[i] = nval[x[i]];
            // cerr << opt[i] << ' ' << l[i] << ' ' << r[i] << ' ' << x[i] << endl;
        }
    }
    for(int i = 0; i < m; ++i){
        // cerr << "#" << i << endl;
        switch(opt[i]){
            case 1: cout << tnt.query(l[i], r[i], x[i] - 1) + 1 << endl; break;
            case 2: cout << rval[tnt.rank(l[i], r[i], rk[i])] << endl; break;
            case 3:{
                tnt.del(ps[i], a[ps[i]]);
                a[ps[i]] = x[i];
                tnt.insert(ps[i], a[ps[i]]);
                break;
            }
            case 4:{
                res = tnt.prev(l[i], r[i], x[i]);
                if(res != -INF)
                    cout << rval[res] << endl;
                else
                    cout << res << endl;
                break;
            }
            case 5:{
                res = tnt.next(l[i], r[i], x[i]);
                if(res != INF)
                    cout << rval[res] << endl;
                else
                    cout << res << endl;
                break;
            }
        }
    }
    return 0;
}
2023/7/19 17:22
加载中...