动态版MKTHNUM问题求助:带更新的区间第k小查询
动态版MKTHNUM问题的持久化线段树更新逻辑问题
我已用持久化线段树解决静态版MKTHNUM问题:给定N个元素的数组,回答查询q(l, r, k),即求子数组l到r排序后的第k小元素。现在要实现数组更新功能——将a[index]赋值为val,同时仍能处理q(l, r, k)类型的查询,但用持久化线段树实现更新逻辑时出错了。
相关代码如下:
#include <iostream> #include <vector> #include <algorithm> using namespace std; const int MAX_VALUE = 1000000; struct Vertex { int sum; Vertex* left, *right; Vertex(int val) : sum(val), left(nullptr), right(nullptr) {} Vertex(Vertex* l, Vertex* r) : sum(0), left(l), right(r) { if (l) sum += l->sum; if (r) sum += r->sum; } }; Vertex* build(int tl, int tr) { if (tl == tr) { return new Vertex(0); } int tm = (tl + tr) / 2; return new Vertex(build(tl, tm), build(tm + 1, tr)); } Vertex* update(Vertex* v, int tl, int tr, int pos) { if (tl == tr) { return new Vertex(v->sum + 1); } int tm = (tl + tr) / 2; if (pos <= tm) { return new Vertex(update(v->left, tl, tm, pos), v->right); } else { return new Vertex(v->left, update(v->right, tm + 1, tr, pos)); } } Vertex* insert(Vertex* old, int tl, int tr, int val) { if (tl == tr) { return new Vertex(old->sum + 1); } int tm = (tl + tr) / 2; if (val <= tm) { return new Vertex(insert(old->left, tl, tm, val), old->right); } else { return new Vertex(old->left, insert(old->right, tm + 1, tr, val)); } } int find_kth(Vertex* nl, Vertex* nr, int tl, int tr, int k) { if (tl == tr) { return tl; } int tm = (tl + tr) / 2, left_count = nr->left->sum - nl->left->sum; if (left_count >= k) { return find_kth(nl->left, nr->left, tl, tm, k); } return find_kth(nl->right, nr->right, tm + 1, tr, k - left_count); } int main() { ios_base::sync_with_stdio(0); cin.tie(0); int n, q; cin >> n; int tl = 0, tr = MAX_VALUE + 1; vector<Vertex*> roots; vector<int> a(n); roots.push_back(build(tl, tr)); for (int j = 0; j < n; j++) cin >> a[j]; for (int i = 0; i < n; i++) { roots.push_back(update(roots.back(), tl, tr, a[i])); } cin >> q; while (q--) { int type, idx, val; cin >> type; if (type == 1) { cin >> idx >> val; a[idx] = val; roots[idx + 1] = insert(roots[idx], tl, tr, val); } if (type == 0) { int x, y, k; cin >> x >> y >> k; int result = find_kth(roots[x], roots[y+1], tl, tr, k+1); cout << result << endl; } } return 0; }
问题根源
你当前的更新逻辑完全错误:
- 静态版持久化线段树是前缀和式版本链,每个
roots[i]代表前i个元素的前缀统计信息。直接修改roots[idx+1]后,后续roots[idx+2]到roots[n]都是基于旧版本构建的,无法反映更新后的数组状态。 insert函数仅做增量插入,而更新操作需要先删除旧值,再插入新值,而非单纯叠加新值。- 前缀式持久化线段树不支持高效单点更新——若强行修复,需从
idx+1到n重新构建所有版本,时间复杂度为O(n log M)(M为值域大小),完全无法应对大数据量。
正确解决方案
动态区间第k小问题,需改用以下高效数据结构:
方案1:树状数组套持久化线段树
- 核心思路:用树状数组维护多个持久化线段树,每个树状数组节点对应值域上的计数统计。
- 单点更新:修改树状数组中对应节点,先删除旧值(delta=-1),再插入新值(delta=1)。
- 区间查询:利用树状数组的前缀和特性,将
r前缀与l-1前缀的线段树做差,在差值树上查找第k小。
方案2:线段树套平衡树
- 核心思路:外层是区间线段树,每个节点用平衡树(如FHQ Treap)存储对应区间的元素,支持排序后的查询操作。
- 单点更新:遍历外层线段树的路径,修改对应平衡树中的元素。
- 区间查询:拆分外层线段树的目标区间,收集所有相关平衡树,合并查询第k小。
推荐代码框架(树状数组套持久化线段树)
#include <iostream> #include <vector> #include <algorithm> using namespace std; const int MAXN = 1e5 + 5; const int MAX_VAL = 1e6 + 5; struct Node { int sum; Node *left, *right; Node(int s=0, Node* l=nullptr, Node* r=nullptr) : sum(s), left(l), right(r) {} }; Node* build(int tl, int tr) { if (tl == tr) return new Node(0); int mid = (tl + tr) / 2; return new Node(0, build(tl, mid), build(mid+1, tr)); } Node* update(Node* prev, int tl, int tr, int pos, int delta) { if (tl == tr) { return new Node(prev->sum + delta); } int mid = (tl + tr) / 2; if (pos <= mid) { return new Node(prev->sum + delta, update(prev->left, tl, mid, pos, delta), prev->right); } else { return new Node(prev->sum + delta, prev->left, update(prev->right, mid+1, tr, pos, delta)); } } vector<Node*> tree[MAXN]; int a[MAXN]; int n, q; void add(int idx, int pos, int delta) { for (; idx <= n; idx += idx & -idx) { tree[idx].push_back(update(tree[idx].back(), 0, MAX_VAL, pos, delta)); } } int query_kth(int l, int r, int k) { vector<Node*> left_nodes, right_nodes; for (int i = l-1; i; i -= i & -i) left_nodes.push_back(tree[i].back()); for (int i = r; i; i -= i & -i) right_nodes.push_back(tree[i].back()); int tl = 0, tr = MAX_VAL; while (tl < tr) { int mid = (tl + tr) / 2; int left_sum = 0; for (int i = 0; i < right_nodes.size(); ++i) { left_sum += right_nodes[i]->left->sum - left_nodes[i]->left->sum; } if (left_sum >= k) { tr = mid; for (auto& node : right_nodes) node = node->left; for (auto& node : left_nodes) node = node->left; } else { tl = mid + 1; k -= left_sum; for (auto& node : right_nodes) node = node->right; for (auto& node : left_nodes) node = node->right; } } return tl; } int main() { ios::sync_with_stdio(false); cin.tie(0); cin >> n; for (int i = 1; i <= n; ++i) { cin >> a[i]; tree[i].push_back(build(0, MAX_VAL)); } for (int i = 1; i <= n; ++i) { add(i, a[i], 1); } cin >> q; while (q--) { int type; cin >> type; if (type == 1) { int idx, val; cin >> idx >> val; add(idx, a[idx], -1); a[idx] = val; add(idx, val, 1); } else { int l, r, k; cin >> l >> r >> k; cout << query_kth(l, r, k) << '\n'; } } return 0; }
内容的提问来源于stack exchange,提问作者José Luis de León
相关产品推荐
相关产品推荐

