初探整体二分

提示:建议开启深色模式。

引入

给定一个序列,要求他的全局第 \(k\) 小,那么可以使用离散化 + 桶然后二分答案就可以求出。但是如果查询区间 \([l,r]\) 的第 \(k\) 小,查询的次数变多,那么解决问题的时间复杂度就变成了 \(O(qn)\),这样显然是会 TLE 的。

对于多次询问,我们发现很多的询问会调用到同一个区间,得到的也是同一个答案。既然如此,我们就考虑将所有的问题集中考虑,这些重复的调用就可以规避掉。

简单介绍

原理解析

下面给定一个序列

\[S = [1,9,1,9,8,1,0] \]

然后我们需要查询区间第 \(k\)

\[Q_1 = [1,7];k_1 = 5\\ Q_2 = [2,5];k_2 = 2\\ Q_3 = [4,7];k_3 = 2 \]

这里以 \(Q_1 = [1,7];k_1 = 5\) 为例!

  • 对于值域 \([0,9]\),先二分求得 \(mid = 4\),左区间为 \([0,4]\),右区间为 \([5,9]\)。树状数组维护的是当前值域左半边(即数值 \(\le mid\)),求得下标在 \([1,7]\) 中,值域为 \([0,4]\) 的数的个数为 \(4\),小于当前的 \(k\),说明 \(Q_1\) 的答案在值域 \([5,9]\) 中。

    询问就变成了

    \[Q_1' = [5,9];k_1' = 5-4 = 1 \]

  • 对于值域 \([5,9]\),先二分求得 \(mid = 7\),左区间为 \([5,7]\),右区间为 \([8,9]\)。使用树状数组查询下标在 \([1,7]\) 中,值域为 \([5,7]\) 的数的个数为 \(0\),小于当前的 \(k\),说明 \(Q_1'\) 的答案在值域 \([8,9]\) 中。

    询问就变成了

    \[Q_1'' = [8,9]; k_1'' = 1-0 = 1 \]

  • 对于值域 \([8,9]\),先二分求得 \(mid = 8\),左区间为 \([8,8]\),右区间为 \([9,9]\)。使用树状数组查询下标在 \([1,7]\) 中,值域在 \([8,8]\) 的数的个数为 \(1\),大于等于当前的 \(k\) 说明 \(Q_1''\) 的答案在值域 \([8,8]\) 中。

    询问就变成了

    \[Q_1''' = [8,8]; k_1''' = 1 \]

  • 现在的询问,它的左端点和右端点相等了,说明询问的答案就是当前的左端点,此时返回答案为 \(8\)

同样的,我们可以画出 \(Q_2\) 的图示。

代码

int n, m;
int a[MAXN], num[MAXN], ch[MAXN];
int ans[MAXM];

struct Opt {    // 使用结构体来存储操作
    int type, l, r, k, id;    // 当 type 为 0 是插入操作,1 为查询操作
} opt[MAXN + MAXM], p1[MAXN + MAXM], p2[MAXN + MAXM];

// l,r 表示操作的下标,L,R 表示查询操作的答案范围值域
void solve(int l, int r, int L, int R) { 
    if (l > r || L > R) return;     // 越界跳过
    if (L == R) {                   // 左右端点重合,得到答案
        for (int i = l; i <= r; i++)
            if (opt[i].type) ans[opt[i].id] = L;
        return;
    }
    int mid = L + R >> 1, cur = 0;  // cur 表示撤销 BIT 中的记录
    int cnt1 = 0, cnt2 = 0;         // p1 和 p2 的计数器
    for (int i = l; i <= r; i++) {
        if (opt[i].type == 0) {     // 如果此时的操作是插入数值的操作
            if (opt[i].k <= mid) {
                bit.add(opt[i].id, 1);  // 树状数组对应位置 +1
                ch[++cur] = opt[i].id;  // 记录,方便撤销操作
                p1[++cnt1] = opt[i];    // 将操作分到左区间
            } else p2[++cnt2] = opt[i]; // 将操作分到右区间
        } else {
            int tmp = bit.query(opt[i].r) - bit.query(opt[i].l - 1); // 查询 [l,mid] 中的数的个数
            if (opt[i].k <= tmp) p1[++cnt1] = opt[i]; // 将操作分到左区间
            else {
                opt[i].k -= tmp;        // 记得将查询的 k 更新
                p2[++cnt2] = opt[i];    // 将操作分到右区间
            }
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i], -1);              // 撤销树状数组的更改
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];         // 将答案值域在左区间的放到一起
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];  // 将答案值域在右区间的放到一起
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R); // 递归求得左区间和右区间
}

一些例题

P3834 可持久化线段树 2(静态区间第 k 小)

双倍经验:P1533 可怜的狗狗

这个就是上面讲的板子。这是整体二分代码中最容易出错的地方。例如 opt[i].k 在插入时代表“数值”,在查询时代表“第 \(k\) 小”;opt[i].id 在插入时代表“位置下标”,在查询时代表“询问编号”。

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int MAXN = 3e5 + 5;
const int MAXM = 5e4 + 5;

int n, m;
int a[MAXN], num[MAXN], ch[MAXN];
int ans[MAXM];

struct Opt {
    int type, l, r, k, id;
} opt[MAXN + MAXM], p1[MAXN + MAXM], p2[MAXN + MAXM];

struct BIT {
    int val[MAXN];
    void add(int x, int k) {
        for (int i = x; i <= n; i += i & (-i)) val[i] += k;
    }
    int query(int x) {
        int res = 0;
        for (int i = x; i; i -= i & (-i)) res += val[i];
        return res;
    }
} bit;

int LSH() {
    for (int i = 1; i <= n; i++) num[i] = a[i];
    sort(num + 1, num + 1 + n);
    int k = unique(num + 1, num + 1 + n) - num - 1;
    for (int i = 1; i <= n; i++) a[i] = lower_bound(num + 1, num + 1 + k, a[i]) - num;
    return k;
}

void solve(int l, int r, int L, int R) {
    if (l > r || L > R) return;
    if (L == R) {
        for (int i = l; i <= r; i++)
            if (opt[i].type) ans[opt[i].id] = L;
        return;
    }
    int mid = L + R >> 1, cur = 0;
    int cnt1 = 0, cnt2 = 0;
    for (int i = l; i <= r; i++) {
        if (opt[i].type == 0) {
            if (opt[i].k <= mid) {
                bit.add(opt[i].id, 1);
                ch[++cur] = opt[i].id;
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else {
            int tmp = bit.query(opt[i].r) - bit.query(opt[i].l - 1);
            if (opt[i].k <= tmp) p1[++cnt1] = opt[i];
            else {
                opt[i].k -= tmp;
                p2[++cnt2] = opt[i];
            }
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i], -1);
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R);
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    //    freopen("file.in","r",stdin);
    //    freopen("file.out","w",stdout);

    cin >> n >> m;
    for (int i = 1; i <= n; i++) cin >> a[i];
    int tot = LSH();

    for (int i = 1; i <= n; i++) opt[i] = {0, 0, 0, a[i], i};
    for (int i = 1; i <= m; i++) {
        int l, r, k;
        cin >> l >> r >> k;
        opt[i + n] = {1, l, r, k, i};
    }
    solve(1, m + n, 1, tot);
    for (int i = 1; i <= m; i++) cout << num[ans[i]] << "\n";

    return 0;
}

P2617 Dynamic Rankings

这个例题就需要你进行修改了,你只需要再加上一种操作:将对应的位置减一即可

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int MAXN = 1e5 + 5;
const int MAXM = 1e5 + 5;
const int MAXOPS = MAXN + MAXM + MAXM;

int n, m;
int a[MAXN], num[MAXOPS];
int ans[MAXM];
pair<int, int> ch[MAXOPS];
char op;
int l, r, k, x, y;

struct Opt {
    int type, l, r, k, id;
} opt[MAXOPS], p1[MAXOPS], p2[MAXOPS], raw[MAXM];

struct BIT {
    int val[MAXN];
    void add(int x, int k) {
        for (int i = x; i <= n; i += i & (-i)) val[i] += k;
    }
    int query(int x) {
        int res = 0;
        for (int i = x; i; i -= i & (-i)) res += val[i];
        return res;
    }
} bit;

int LSH(int vcnt) {
    sort(num + 1, num + 1 + vcnt);
    int k = unique(num + 1, num + 1 + vcnt) - num - 1;
    for (int i = 1; i <= n; i++) a[i] = lower_bound(num + 1, num + 1 + k, a[i]) - num;
    return k;
}

void solve(int l, int r, int L, int R) {
    if (l > r || L > R) return;
    if (L == R) {
        for (int i = l; i <= r; i++)
            if (opt[i].type == 2) ans[opt[i].id] = L;
        return;
    }
    int mid = L + R >> 1, cur = 0;
    int cnt1 = 0, cnt2 = 0;
    for (int i = l; i <= r; i++) {
        if (opt[i].type == 0) {             // 插入操作
            if (opt[i].k <= mid) {
                bit.add(opt[i].id, 1);
                ch[++cur] = {opt[i].id, 1};
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else if (opt[i].type == 1) {      // 删除操作
            if (opt[i].k <= mid) {
                bit.add(opt[i].id, -1);
                ch[++cur] = {opt[i].id, -1};
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else {
            int tmp = bit.query(opt[i].r) - bit.query(opt[i].l - 1);
            if (opt[i].k <= tmp) p1[++cnt1] = opt[i];
            else {
                opt[i].k -= tmp;
                p2[++cnt2] = opt[i];
            }
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i].first, -ch[i].second);
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R);
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    //	freopen("file.in","r",stdin);
    //	freopen("file.out","w",stdout);

    cin >> n >> m;
    for (int i = 1; i <= n; i++) cin >> a[i];

    int cnt = 0, tot = 0, vcnt = n;
    for (int i = 1; i <= n; i++) num[i] = a[i];
    for (int i = 1; i <= m; i++) {
        cin >> op;
        if (op == 'Q') {
            cin >> l >> r >> k;
            raw[++cnt] = {2, l, r, k, ++tot};
        }
        if (op == 'C') {
            cin >> x >> y;
            raw[++cnt] = {1, 0, 0, y, x};
            num[++vcnt] = y;
        }
    }

    int w = LSH(vcnt);

    int cc = 0;
    for (int i = 1; i <= n; i++) opt[++cc] = {0, 0, 0, a[i], i};
    for (int i = 1; i <= cnt; i++) {
        if (raw[i].type == 2) {
            opt[++cc] = raw[i];
        } else {
            int pos = raw[i].id;
            int nv = lower_bound(num + 1, num + 1 + w, raw[i].k) - num;
            opt[++cc] = {1, 0, 0, a[pos], pos}; //记得删除后再插入
            opt[++cc] = {0, 0, 0, nv, pos};
            a[pos] = nv;
        }
    }

    solve(1, cc, 1, w);
    for (int i = 1; i <= tot; i++) cout << num[ans[i]] << "\n";
    return 0;
}

P1527 [国家集训队] 矩阵乘法

这里从一维变成了二维,就需要二维的 BIT 来维护数据了。

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int MAXN = 505;
const int MAXQ = 4e5 + 5;

int n, q, a[MAXN][MAXN];
int num[250005], ans[MAXQ];
pair<int, int> ch[250005];

struct Opt {
    int type, x1, x2, y1, y2, k, id1, id2;
} opt[MAXQ], p1[MAXQ], p2[MAXQ];

struct BIT {
    int val[MAXN][MAXN];

    BIT() {
        memset(val, 0, sizeof val);
    }

    void add(int x, int y, int k) {
        for (int i = x; i <= n; i += i & (-i))
            for (int j = y; j <= n; j += j & (-j)) val[i][j] += k;
    }

    int query(int x, int y) {
        int res = 0;
        for (int i = x; i; i -= i & (-i))
            for (int j = y; j; j -= (j & -j)) res += val[i][j];
        return res;
    }

} bit;

int LSH() {
    int tot = 0;
    for (int i = 1; i <= n; i++)
        for (int j = 1; j <= n; j++) num[++tot] = a[i][j];
    sort(num + 1, num + 1 + tot);
    int k = unique(num + 1, num + 1 + tot) - num - 1;
    for (int i = 1; i <= n; i++)
        for (int j = 1; j <= n; j++) a[i][j] = lower_bound(num + 1, num + 1 + k, a[i][j]) - num;
    return k;
}

void solve(int l, int r, int L, int R) {
    if (l > r || L > R) return;
    if (L == R) {
        for (int i = l; i <= r; i++)
            if (opt[i].type) ans[opt[i].id1] = L;
        return;
    }
    int mid = L + R >> 1, cur = 0;
    int cnt1 = 0, cnt2 = 0;
    for (int i = l; i <= r; i++) {
        if (!opt[i].type) {
            if (opt[i].k <= mid) {
                bit.add(opt[i].id1, opt[i].id2, 1);
                ch[++cur] = {opt[i].id1, opt[i].id2};
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else {
            int tmp = bit.query(opt[i].x2, opt[i].y2) - bit.query(opt[i].x1 - 1, opt[i].y2) - bit.query(opt[i].x2, opt[i].y1 - 1) + bit.query(opt[i].x1 - 1, opt[i].y1 - 1);
            if (opt[i].k <= tmp) p1[++cnt1] = opt[i];
            else {
                opt[i].k -= tmp;
                p2[++cnt2] = opt[i];
            }
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i].first, ch[i].second, -1);
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R);
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    //	freopen("file.in","r",stdin);
    //	freopen("file.out","w",stdout);
    cin >> n >> q;
    for (int i = 1; i <= n; i++)
        for (int j = 1; j <= n; j++) cin >> a[i][j];

    int tot = LSH(), cnt = 0;
    for (int i = 1; i <= n; i++)
        for (int j = 1; j <= n; j++) opt[++cnt] = {0, 0, 0, 0, 0, a[i][j], i, j};

    for (int i = 1; i <= q; i++) {
        int x1, x2, y1, y2, k;
        cin >> x1 >> y1 >> x2 >> y2 >> k;
        opt[++cnt] = {1, x1, x2, y1, y2, k, i, 0};
    }
    solve(1, cnt, 1, tot);
    for (int i = 1; i <= q; i++) cout << num[ans[i]] << "\n";

    return 0;
}

P3527 [POI 2011] MET-Meteors

这个题目涉及到了区间加和,单点查询,所以考虑差分来维护。

#include <bits/stdc++.h>
using namespace std;

const int MAXN = 3e5 + 5;

int n, m, q;
int o[MAXN], p[MAXN], ans[MAXN];
pair<int, int> ch[MAXN << 1];
vector<int> G[MAXN];            // 存储国家对应空间站的位置

struct Opt {
    int type, l, r, k, id, val;
} opt[MAXN << 1], p1[MAXN << 1], p2[MAXN << 1];

struct BIT {
    unsigned long long val[MAXN << 1];

    BIT() {
        memset(val, 0, sizeof val);
    }

    void add(int x, int k) {
        for (int i = x; i <= m; i += i & (-i)) val[i] += k;
    }

    unsigned long long query(int x) {
        unsigned long long res = 0;
        for (int i = x; i; i -= i & (-i)) res += val[i];
        return res;
    }

} bit;

void solve(int l, int r, int L, int R) {
    if (l > r || L > R) return;
    if (L == R) {
        for (int i = l; i <= r; i++)
            if (opt[i].type) ans[opt[i].id] = L;
        return;
    }
    int mid = (L + R) >> 1, cur = 0;
    int cnt1 = 0, cnt2 = 0;
    for (int i = l; i <= r; i++) {
        if (!opt[i].type) {
            if (opt[i].k <= mid) {
                if (opt[i].l <= opt[i].r) {
                    bit.add(opt[i].l, opt[i].val), bit.add(opt[i].r + 1, -opt[i].val);
                    ch[++cur] = {opt[i].l, opt[i].val};
                    ch[++cur] = {opt[i].r + 1, -opt[i].val};
                } else {
                    bit.add(opt[i].l, opt[i].val), bit.add(m + 1, -opt[i].val);
                    bit.add(1, opt[i].val), bit.add(opt[i].r + 1, -opt[i].val);
                    ch[++cur] = {opt[i].l, opt[i].val};
                    ch[++cur] = {m + 1, -opt[i].val};
                    ch[++cur] = {1, opt[i].val};
                    ch[++cur] = {opt[i].r + 1, -opt[i].val};
                }
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else {
            unsigned long long tmp = 0;
            for (int j : G[opt[i].k]) tmp += bit.query(j);
            if (p[opt[i].k] <= tmp) p1[++cnt1] = opt[i];
            else p[opt[i].k] -= tmp, p2[++cnt2] = opt[i];
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i].first, -ch[i].second);
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R);
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    //	freopen("file.in","r",stdin);
    //	freopen("file.out","w",stdout);
    cin >> n >> m;
    for (int i = 1; i <= m; i++) cin >> o[i];
    for (int i = 1; i <= n; i++) cin >> p[i];
    for (int i = 1; i <= m; i++) G[o[i]].push_back(i);

    cin >> q;
    for (int i = 1; i <= q; i++) {
        int l, r, a;
        cin >> l >> r >> a;
        opt[i] = {0, l, r, i, 0, a};
    }
    for (int i = 1; i <= n; i++) opt[i + q] = {1, 0, 0, i, i, 0};
    solve(1, n + q, 1, q + 1);
    for (int i = 1; i <= n; i++) {
        if (ans[i] <= q) cout << ans[i] << "\n";
        else cout << "NIE" << "\n";
    }

    return 0;
}

P7424 [THUPC 2017] 天天爱射击

这个题目只需考虑一些简单的反演:把每个子弹能击穿的木板数转换成每块木板被哪一个子弹击穿的。然后把每一个子弹的位置视作下标,把每一个子弹的射出顺序视作值。对于每一块木板,只需要查询 \([x_1,x_2]\) 中的第 \(s\) 小的值,最后将查询到的子弹的结果加一。

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int MAXN = 2e5 + 5;
const int MAXM = 2e5 + 5;

int n, m;
int a, num[MAXN], ch[MAXN];
int ans[MAXM], cnt[MAXN];

struct Opt {
    int type, l, r, k, id;
} opt[MAXN + MAXM], p1[MAXN + MAXM], p2[MAXN + MAXM];

struct BIT {
    int val[MAXN];
    void add(int x, int k) {
        for (int i = x; i <= 2e5; i += i & (-i)) val[i] += k;
    }
    int query(int x) {
        int res = 0;
        for (int i = x; i; i -= i & (-i)) res += val[i];
        return res;
    }
} bit;

void solve(int l, int r, int L, int R) {
    if (l > r || L > R) return;
    if (L == R) {
        for (int i = l; i <= r; i++)
            if (opt[i].type) ans[opt[i].id] = L;
        return;
    }
    int mid = L + R >> 1, cur = 0;
    int cnt1 = 0, cnt2 = 0;
    for (int i = l; i <= r; i++) {
        if (opt[i].type == 0) {
            if (opt[i].k <= mid) {
                bit.add(opt[i].id, 1);
                ch[++cur] = opt[i].id;
                p1[++cnt1] = opt[i];
            } else p2[++cnt2] = opt[i];
        } else {
            int tmp = bit.query(opt[i].r) - bit.query(opt[i].l - 1);
            if (opt[i].k <= tmp) p1[++cnt1] = opt[i];
            else {
                opt[i].k -= tmp;
                p2[++cnt2] = opt[i];
            }
        }
    }
    for (int i = 1; i <= cur; i++) bit.add(ch[i], -1);
    for (int i = 1; i <= cnt1; i++) opt[l + i - 1] = p1[i];
    for (int i = 1; i <= cnt2; i++) opt[l + cnt1 + i - 1] = p2[i];
    solve(l, l + cnt1 - 1, L, mid), solve(l + cnt1, r, mid + 1, R);
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    //	freopen("file.in","r",stdin);
    //	freopen("file.out","w",stdout);

    cin >> n >> m;

    int tot = 0;
    for (int i = 1; i <= n; i++) {
        int l, r, s;
        cin >> l >> r >> s;
        opt[m + i] = {1, l, r, s, i};
    }
    for (int i = 1; i <= m; i++) {
        cin >> a;
        opt[++tot] = {0, 0, 0, i, a};
    }
    solve(1, n + m, 1, m + 1);
    for (int i = 1; i <= n; i++) cnt[ans[i]]++;
    for (int i = 1; i <= m; i++) cout << cnt[i] << "\n";

    return 0;
}

P3332 [ZJOI2013] K 大数查询

这个涉及到了区间修改、区间查询,考虑使用线段树进行维护。

posted @ 2026-08-30 21:34  C0nfidence  阅读(1)  评论(0)    收藏  举报