如何卡常/kk
查看原帖
如何卡常/kk
235671
fireinice楼主2022/5/31 19:54

目前94分,#1#2#3 TLE,写的是时间分块+序列分块,不知哪里还能再优化,已经试过有:

  • 调块长
  • 预处理块左右端点,长度为i的贡献gx[i]
  • register,O2,快读快写
  • 手写栈,邻接表

PS:为避免无用代码过长去掉了快读快写

#include <algorithm>
#include <cmath>
#include <cstdio>
#include <cstring>
#include <iostream>
#include <stack>

using namespace std;
const int N = 3e5 + 500, SQ = N, SN = N;
typedef long long ll;
#define rint register int
#define rll register long long
#define rbool register bool

struct Upd {
    int idx, v, id;
    bool operator==(Upd b) { return idx == b.idx; }
} upds[N];
int ucnt = 0;

struct Query {
    int x, l, r, time, id;
} qs[N];
ll anses[SN];
int qcnt = 0;

int n, m;
int a[N];
struct SeqBlock {
    int b[N];
    ll ans[SN];
    int pos[N];
    int sn;

    struct Opt {
        int idx;
        ll det;
        int bg, ed;
    };

    Opt stk[N];
    int top;
    // int lft(int x) { return (x - 1) * sn + 1; }
    // int blk(int x) { return (x - 1) / sn + 1; }

    int lft[SN], blk[N];
    ll gx[N];
    
#define lft(x) lft[x]
#define blk(x) blk[x]

    void init() {
        for (int i = 1; i <= n + 10; i++) {
            blk[i] = (i - 1) / sn + 1;
        }
        for (int i = 1; i <= n / sn + 10; i++) {
            lft[i] = (i - 1) * sn + 1;
        }
        for (int i = 1; i <= n; i++) {
            gx[i] = (ll)i * (i + 1) >> 1;
        }
        top = 0;
    }

    void build() {
        memset(ans, 0, sizeof(ans[0]) * (n / sn + 5));
        memset(b, 0, sizeof(b));
        memset(pos, 0, sizeof(b));
        top = 0;
    }

#define calc(len) gx[len]

    inline void update(rint idx, bool willundo = true) {
        rint bk = blk(idx), last = idx - 1, next = idx + 1, ld = lft(bk),
             rd = lft(bk + 1) - 1;
        rbool lone = b[last] && blk(last) == bk,
              rone = b[next] && blk(next) == bk;
        rbool lc = pos[ld] == last || idx == ld,
              rc = pos[rd] == next || idx == rd;
        rll det = 0;
        rint ed = pos[next], bg = pos[last];
        if (lone && rone) {
            if (!(lc && rc)) {
                if (lc) {
                    det = -calc(ed - next + 1);
                } else if (rc) {
                    det = -calc(last - bg + 1);
                } else {
                    det = ll(ed - next + 2) * (last - bg + 2);
                }
            }
            pos[last] = pos[next] = 0;
            pos[ed] = bg;
            pos[bg] = ed;
        } else if (lone) {
            if (rc)
                det = -calc(last - bg + 1);
            else if (!lc)
                det = (last - bg + 2);
            pos[idx] = bg;
            pos[bg] = idx;
        } else if (rone) {
            if (lc)
                det = -calc(ed - next + 1);
            else if (!rc)
                det = (ed - next + 2);
            pos[idx] = ed;
            pos[ed] = idx;
        } else {
            pos[idx] = idx;
            if (!lc && !rc)
                det = 1;
        }
        
        if (willundo)
            stk[++top] = ((Opt){idx, det, bg, ed});

        ans[bk] += det;
        b[idx] = true;
    }

    void undo() {
        register Opt &opt = stk[top];
        top--;
        rint bk = blk(opt.idx);
        rint idx = opt.idx;
        b[idx] = false;
        ans[bk] -= opt.det;
        rint last = idx - 1, next = idx + 1, ld = lft(bk), rd = lft(bk + 1) - 1;
        rbool lone = b[last] && blk(last) == bk,
              rone = b[next] && blk(next) == bk;
        if (lone && rone) {
            pos[opt.ed] = idx + 1;
            pos[opt.bg] = idx - 1;
            pos[idx - 1] = opt.bg;
            pos[idx + 1] = opt.ed;
        } else if (lone) {
            pos[opt.bg] = idx - 1;
            pos[idx - 1] = opt.bg;
        } else if (rone) {
            pos[opt.ed] = idx + 1;
            pos[idx + 1] = opt.ed;
        }
        pos[idx] = 0;
    }

    ll scan(int l,
            int r,
            int tag,
            int last = 0) {  // tag:1->全暴力2->左散块3->右散块
        rll res = 0;
        rint bk = blk(l);
        for (rint i = l; i <= r; i++) {
            if (b[i] && !last)
                last = i;
            if (!b[i] && last) {
                if (tag == 2 && b[lft(bk + 1) - 1] &&
                    last >= pos[lft(bk + 1) - 1])
                    break;
                int len = i - last;
                res += calc(len);
                last = 0;
            }
        }
        if (tag != 2 && b[r]) {
            res += calc(r - last + 1);
        }
        return res;
    }

    ll query(int l, int r) {
        int bl = blk(l), br = blk(r);
        if (br - bl <= 1)
            return scan(l, r, 1);
        bl++, br--;
        rll res = 0;
        res += scan(l, lft(bl) - 1, 2);
        rint last = max(b[lft(bl) - 1] ? pos[lft(bl) - 1] : lft(bl), l);
        for (rint i = bl; i <= br; i++) {
            if (pos[lft(i)] == lft(i + 1) - 1) {
                continue;
            } else {
                rint len = (b[lft(i)] ? pos[lft(i)] : lft(i) - 1) - last + 1;
                res += calc(len);
                last = b[lft(i + 1) - 1] ? pos[lft(i + 1) - 1] : lft(i + 1);
                res += ans[i];
            }
        }
        res += scan(lft(br + 1), r, 3, last);
        return res;
    }

} seq;

int qcmp(const Query& a, const Query& b) {
    return a.x < b.x;
}

bool ignored[N];

struct Edge {
    int v, nxt;
} edges[N];
int head[N], ecnt = 0;
void add_edge(int u, int v) {
    edges[++ecnt] = (Edge){v, head[u]};
    head[u] = ecnt;
}
bool changed[N];

void solve() {
    seq.build();
    sort(qs + 1, qs + 1 + qcnt, qcmp);
    for (int i = 1; i <= ucnt; i++) {
        ignored[upds[i].idx] = true;
    }
    memset(head, 0, sizeof(head));
    ecnt = 0;
    for (int i = 1; i <= n; i++) {
        if (!ignored[i])
            add_edge(a[i], i);
    }
    int p = 1;

    for (int i = 1; i <= qcnt; i++) {
        for (; p <= n && p <= qs[i].x; p++) {
            for (int k = head[p]; k; k = edges[k].nxt) {
                seq.update(edges[k].v, false);
            }
        }
        int cnt = 0;
        for (int j = qs[i].time; j >= 1; j--) {
            if (!changed[upds[j].idx] && upds[j].v <= qs[i].x) {
                seq.update(upds[j].idx);
            }
            changed[upds[j].idx] = true;
        }
        for (int j = qs[i].time + 1; j <= ucnt; j++) {
            if (!changed[upds[j].idx] && a[upds[j].idx] <= qs[i].x) {
                seq.update(upds[j].idx);
            }
            changed[upds[j].idx] = true;
        }
        anses[qs[i].id] = seq.query(qs[i].l, qs[i].r);
        while (seq.top)
            seq.undo();
        for (int j = 1; j <= ucnt; j++) {
            changed[upds[j].idx] = false;
        }
    }

    for (int i = 1; i <= qcnt; i++) {
        println(anses[i]);
    }

    for (int i = 1; i <= ucnt; i++) {
        a[upds[i].idx] = upds[i].v;
        ignored[upds[i].idx] = false;
    }
}

signed main() {
    read(n);
    read(m);
    seq.sn = 500;
    int qn = 1000;
    seq.init();

    for (int i = 1; i <= n; i++) {
        read(a[i]);
    }
    seq.build();
    int t = 1;
    while (t <= m) {
        ucnt = qcnt = 0;
        for (; t <= m && ucnt <= qn; t++) {
            int op, l, r, x, y;
            read(op);
            if (op == 1) {
                read(x);
                read(y);
                upds[++ucnt] = (Upd){x, y, ucnt};
            } else {
                read(l);
                read(r);
                read(x);
                qs[++qcnt] = (Query){x, l, r, ucnt, qcnt};
            }
        }
        solve();
    }
}

2022/5/31 19:54
加载中...