fork download
  1. #include <bits/stdc++.h>
  2. using namespace std;
  3. using ll = long long;
  4.  
  5. // ===================================================================
  6. // PERSISTENT SEGMENT TREE LIBRARY
  7. //
  8. // Contains:
  9. // 1) Persistent Segment Tree with lazy propagation
  10. // - Range add / range sum queries
  11. // - Also provides "first index with prefix sum > target"
  12. // 2) Persistent Segment Tree for range assignment (set)
  13. // - Range set / range sum queries
  14. // 3) Persistent Segment Tree for point updates and k‑th smallest
  15. // - (no lazy, count‑based, classic for subarray order statistics)
  16. //
  17. // All classes are self‑contained, ready to use as black boxes.
  18. // Time complexities: O(log N) per operation, O(N + Q log N) memory.
  19. // ===================================================================
  20.  
  21. // ===================================================================
  22. // 1) PERSISTENT SEGMENT TREE – RANGE ADD, RANGE SUM
  23. // ===================================================================
  24.  
  25. class PersistentSegTreeLazy {
  26. private:
  27. struct Node {
  28. int left = 0, right = 0;
  29. ll sum = 0, lazy = 0;
  30. Node() {}
  31. Node(int l, int r, ll s, ll la) : left(l), right(r), sum(s), lazy(la) {}
  32. };
  33.  
  34. vector<Node> tree; // node pool (index 0 = dummy)
  35. int N; // size of the underlying array
  36. vector<int> roots; // roots[version] = root node index
  37.  
  38. // Creates a copy of an existing node. Returns the new node's index.
  39. int cloneNode(int node) {
  40. tree.push_back(tree[node]);
  41. return (int)tree.size() - 1;
  42. }
  43.  
  44. // Applies an addition to a node's whole segment (persistent – call on a cloned node).
  45. void apply(int node, int l, int r, ll val) {
  46. tree[node].sum += val * (r - l + 1);
  47. tree[node].lazy += val;
  48. }
  49.  
  50. // Pushes the lazy value down to children (clones children to keep persistence).
  51. int push(int node, int l, int r) {
  52. if (l == r || tree[node].lazy == 0) return node;
  53. int mid = (l + r) / 2;
  54. int leftChild = cloneNode(tree[node].left);
  55. int rightChild = cloneNode(tree[node].right);
  56. apply(leftChild, l, mid, tree[node].lazy);
  57. apply(rightChild, mid + 1, r, tree[node].lazy);
  58. tree[node].left = leftChild;
  59. tree[node].right = rightChild;
  60. tree[node].lazy = 0;
  61. return node;
  62. }
  63.  
  64. // Builds the initial tree from the array.
  65. int build(const vector<ll>& arr, int l, int r) {
  66. int node = (int)tree.size();
  67. tree.push_back(Node());
  68. if (l == r) {
  69. tree[node].sum = arr[l];
  70. return node;
  71. }
  72. int mid = (l + r) / 2;
  73. int leftChild = build(arr, l, mid);
  74. int rightChild = build(arr, mid + 1, r);
  75. tree[node].left = leftChild;
  76. tree[node].right = rightChild;
  77. tree[node].sum = tree[leftChild].sum + tree[rightChild].sum;
  78. return node;
  79. }
  80.  
  81. // Recursive range add – returns new root after the update.
  82. int updateRangeAddRec(int node, int l, int r, int ql, int qr, ll val) {
  83. int newNode = cloneNode(node);
  84. if (ql <= l && r <= qr) {
  85. apply(newNode, l, r, val);
  86. return newNode;
  87. }
  88. newNode = push(newNode, l, r);
  89. int mid = (l + r) / 2;
  90. if (ql <= mid) {
  91. int newLeft = updateRangeAddRec(tree[newNode].left, l, mid, ql, qr, val);
  92. tree[newNode].left = newLeft;
  93. }
  94. if (qr > mid) {
  95. int newRight = updateRangeAddRec(tree[newNode].right, mid + 1, r, ql, qr, val);
  96. tree[newNode].right = newRight;
  97. }
  98. tree[newNode].sum = tree[tree[newNode].left].sum + tree[tree[newNode].right].sum;
  99. return newNode;
  100. }
  101.  
  102. // Recursive range sum – does NOT modify the tree.
  103. ll queryRangeSumRec(int node, int l, int r, int ql, int qr) {
  104. if (ql <= l && r <= qr) {
  105. return tree[node].sum;
  106. }
  107. int mid = (l + r) / 2;
  108. ll res = 0;
  109. if (ql <= mid) {
  110. res += queryRangeSumRec(tree[node].left, l, mid, ql, qr);
  111. }
  112. if (qr > mid) {
  113. res += queryRangeSumRec(tree[node].right, mid + 1, r, ql, qr);
  114. }
  115. // Add the contribution of the lazy value stored in this node.
  116. int overlapL = max(ql, l);
  117. int overlapR = min(qr, r);
  118. if (overlapL <= overlapR) {
  119. res += tree[node].lazy * (overlapR - overlapL + 1);
  120. }
  121. return res;
  122. }
  123.  
  124. // Recursive search for the first index where prefix sum > target.
  125. // accLazy = sum of lazy values from ancestors (not yet applied to this node).
  126. int findFirstPrefixGreaterRec(int node, int l, int r, ll target, ll accLazy) {
  127. // Total sum of this segment if ancestor lazy were applied.
  128. ll total = tree[node].sum + accLazy * (r - l + 1);
  129. if (total <= target) return -1;
  130. if (l == r) return l;
  131.  
  132. int mid = (l + r) / 2;
  133. ll newAccLazy = accLazy + tree[node].lazy; // passed to children
  134. int leftChild = tree[node].left;
  135. ll leftSum = tree[leftChild].sum + newAccLazy * (mid - l + 1);
  136.  
  137. if (leftSum > target) {
  138. return findFirstPrefixGreaterRec(leftChild, l, mid, target, newAccLazy);
  139. } else {
  140. return findFirstPrefixGreaterRec(tree[node].right, mid + 1, r,
  141. target - leftSum, newAccLazy);
  142. }
  143. }
  144.  
  145. public:
  146. // Constructor: builds the tree from the initial array.
  147. // Version 0 is created automatically.
  148. PersistentSegTreeLazy(const vector<ll>& arr) {
  149. N = (int)arr.size();
  150. tree.clear();
  151. tree.reserve(2000000); // adjust size for your problem
  152. tree.push_back(Node()); // dummy node at index 0
  153. int root = build(arr, 0, N - 1);
  154. roots.clear();
  155. roots.push_back(root);
  156. }
  157.  
  158. // Returns the root index of a given version.
  159. int getRoot(int version) const {
  160. return roots[version];
  161. }
  162.  
  163. // Returns the number of the latest version.
  164. int getCurrentVersion() const {
  165. return (int)roots.size() - 1;
  166. }
  167.  
  168. // Stores a new root as a new version.
  169. void addVersion(int root) {
  170. roots.push_back(root);
  171. }
  172.  
  173. // Applies a range add to the given version root.
  174. // Returns the new root (caller should use addVersion() to save it).
  175. int updateRangeAdd(int root, int ql, int qr, ll val) {
  176. return updateRangeAddRec(root, 0, N - 1, ql, qr, val);
  177. }
  178.  
  179. // Queries the sum on [ql, qr] for the given version root.
  180. ll queryRangeSum(int root, int ql, int qr) {
  181. return queryRangeSumRec(root, 0, N - 1, ql, qr);
  182. }
  183.  
  184. // Finds the smallest index i (0‑based) such that sum(a[0..i]) > target.
  185. // If none, returns -1.
  186. int findFirstPrefixGreater(int root, ll target) {
  187. return findFirstPrefixGreaterRec(root, 0, N - 1, target, 0);
  188. }
  189. };
  190.  
  191. // ===================================================================
  192. // 2) PERSISTENT SEGMENT TREE – RANGE ASSIGNMENT (SET), RANGE SUM
  193. // ===================================================================
  194.  
  195. class PersistentSegTreeSet {
  196. private:
  197. struct Node {
  198. int left = 0, right = 0;
  199. ll sum = 0;
  200. ll lazy = 0;
  201. bool hasLazy = false; // true if there is a pending assignment
  202. Node() {}
  203. Node(int l, int r, ll s, ll la, bool hl)
  204. : left(l), right(r), sum(s), lazy(la), hasLazy(hl) {}
  205. };
  206.  
  207. vector<Node> tree;
  208. int N;
  209. vector<int> roots;
  210.  
  211. int cloneNode(int node) {
  212. tree.push_back(tree[node]);
  213. return (int)tree.size() - 1;
  214. }
  215.  
  216. void applySet(int node, int l, int r, ll val) {
  217. tree[node].sum = val * (r - l + 1);
  218. tree[node].lazy = val;
  219. tree[node].hasLazy = true;
  220. }
  221.  
  222. int push(int node, int l, int r) {
  223. if (l == r || !tree[node].hasLazy) return node;
  224. int mid = (l + r) / 2;
  225. int leftChild = cloneNode(tree[node].left);
  226. int rightChild = cloneNode(tree[node].right);
  227. applySet(leftChild, l, mid, tree[node].lazy);
  228. applySet(rightChild, mid + 1, r, tree[node].lazy);
  229. tree[node].left = leftChild;
  230. tree[node].right = rightChild;
  231. tree[node].hasLazy = false;
  232. return node;
  233. }
  234.  
  235. int build(const vector<ll>& arr, int l, int r) {
  236. int node = (int)tree.size();
  237. tree.push_back(Node());
  238. if (l == r) {
  239. tree[node].sum = arr[l];
  240. return node;
  241. }
  242. int mid = (l + r) / 2;
  243. int leftChild = build(arr, l, mid);
  244. int rightChild = build(arr, mid + 1, r);
  245. tree[node].left = leftChild;
  246. tree[node].right = rightChild;
  247. tree[node].sum = tree[leftChild].sum + tree[rightChild].sum;
  248. return node;
  249. }
  250.  
  251. int updateRangeSetRec(int node, int l, int r, int ql, int qr, ll val) {
  252. int newNode = cloneNode(node);
  253. if (ql <= l && r <= qr) {
  254. applySet(newNode, l, r, val);
  255. return newNode;
  256. }
  257. newNode = push(newNode, l, r);
  258. int mid = (l + r) / 2;
  259. if (ql <= mid) {
  260. int newLeft = updateRangeSetRec(tree[newNode].left, l, mid, ql, qr, val);
  261. tree[newNode].left = newLeft;
  262. }
  263. if (qr > mid) {
  264. int newRight = updateRangeSetRec(tree[newNode].right, mid + 1, r, ql, qr, val);
  265. tree[newNode].right = newRight;
  266. }
  267. tree[newNode].sum = tree[tree[newNode].left].sum + tree[tree[newNode].right].sum;
  268. return newNode;
  269. }
  270.  
  271. ll queryRangeSumRec(int node, int l, int r, int ql, int qr) {
  272. if (ql <= l && r <= qr) {
  273. return tree[node].sum;
  274. }
  275. int mid = (l + r) / 2;
  276. ll res = 0;
  277. if (ql <= mid)
  278. res += queryRangeSumRec(tree[node].left, l, mid, ql, qr);
  279. if (qr > mid)
  280. res += queryRangeSumRec(tree[node].right, mid + 1, r, ql, qr);
  281.  
  282. // If this node has a pending assignment, it covers the whole segment.
  283. int overlapL = max(ql, l);
  284. int overlapR = min(qr, r);
  285. if (overlapL <= overlapR && tree[node].hasLazy) {
  286. res += tree[node].lazy * (overlapR - overlapL + 1);
  287. }
  288. return res;
  289. }
  290.  
  291. public:
  292. PersistentSegTreeSet(const vector<ll>& arr) {
  293. N = (int)arr.size();
  294. tree.clear();
  295. tree.reserve(2000000);
  296. tree.push_back(Node());
  297. int root = build(arr, 0, N - 1);
  298. roots.clear();
  299. roots.push_back(root);
  300. }
  301.  
  302. int getRoot(int version) const { return roots[version]; }
  303. int getCurrentVersion() const { return (int)roots.size() - 1; }
  304. void addVersion(int root) { roots.push_back(root); }
  305.  
  306. // Range assignment: sets every element in [ql, qr] to val.
  307. int updateRangeSet(int root, int ql, int qr, ll val) {
  308. return updateRangeSetRec(root, 0, N - 1, ql, qr, val);
  309. }
  310.  
  311. // Range sum query.
  312. ll queryRangeSum(int root, int ql, int qr) {
  313. return queryRangeSumRec(root, 0, N - 1, ql, qr);
  314. }
  315. };
  316.  
  317. // ===================================================================
  318. // 3) PERSISTENT SEGMENT TREE FOR K‑TH SMALLEST (COUNT TREE)
  319. // No lazy – point updates, frequency based.
  320. // ===================================================================
  321.  
  322. class PersistentSegTreeCount {
  323. private:
  324. struct Node {
  325. int left = 0, right = 0;
  326. int cnt = 0; // number of elements in this segment
  327. Node() {}
  328. Node(int l, int r, int c) : left(l), right(r), cnt(c) {}
  329. };
  330.  
  331. vector<Node> tree;
  332. int N; // number of distinct values (coordinates)
  333. vector<int> roots; // roots[0] = empty tree, roots[i] = after first i elements
  334.  
  335. int build(int l, int r) {
  336. int node = (int)tree.size();
  337. tree.push_back(Node());
  338. if (l == r) return node;
  339. int mid = (l + r) / 2;
  340. int leftChild = build(l, mid);
  341. int rightChild = build(mid + 1, r);
  342. tree[node].left = leftChild;
  343. tree[node].right = rightChild;
  344. return node;
  345. }
  346.  
  347. int updatePointRec(int node, int l, int r, int pos, int delta) {
  348. int newNode = (int)tree.size();
  349. tree.push_back(tree[node]);
  350. tree[newNode].cnt += delta;
  351. if (l == r) return newNode;
  352. int mid = (l + r) / 2;
  353. if (pos <= mid) {
  354. int newLeft = updatePointRec(tree[node].left, l, mid, pos, delta);
  355. tree[newNode].left = newLeft;
  356. } else {
  357. int newRight = updatePointRec(tree[node].right, mid + 1, r, pos, delta);
  358. tree[newNode].right = newRight;
  359. }
  360. return newNode;
  361. }
  362.  
  363. // k is 1‑indexed. nodeL = earlier root, nodeR = later root.
  364. int queryKthRec(int nodeL, int nodeR, int l, int r, int k) {
  365. if (l == r) return l;
  366. int mid = (l + r) / 2;
  367. int leftCount = tree[tree[nodeR].left].cnt - tree[tree[nodeL].left].cnt;
  368. if (k <= leftCount) {
  369. return queryKthRec(tree[nodeL].left, tree[nodeR].left, l, mid, k);
  370. } else {
  371. return queryKthRec(tree[nodeL].right, tree[nodeR].right,
  372. mid + 1, r, k - leftCount);
  373. }
  374. }
  375.  
  376. public:
  377. // distinctValues = number of different compressed coordinates.
  378. PersistentSegTreeCount(int distinctValues) {
  379. N = distinctValues;
  380. tree.clear();
  381. tree.reserve(2000000);
  382. tree.push_back(Node()); // dummy node 0
  383. int root = build(1, N); // 1‑indexed for easier handling
  384. roots.clear();
  385. roots.push_back(root); // version 0 : empty
  386. }
  387.  
  388. // Returns the root of the given version.
  389. int getRoot(int version) const {
  390. return roots[version];
  391. }
  392.  
  393. // Adds a new version by applying a point update (+delta) to the previous root.
  394. int updatePoint(int previousRoot, int pos, int delta) {
  395. return updatePointRec(previousRoot, 1, N, pos, delta);
  396. }
  397.  
  398. // Appends a new version (root) to the internal list.
  399. void addVersion(int root) {
  400. roots.push_back(root);
  401. }
  402.  
  403. // Query the k‑th smallest (1‑indexed) in the subarray [l, r] (0‑based indices).
  404. // Version indices correspond to prefix lengths:
  405. // versionL = l (prefix before the subarray)
  406. // versionR = r + 1 (prefix up to r)
  407. // For example, to query a[l..r], call queryKth(l, r+1, k).
  408. int queryKth(int versionL, int versionR, int k) {
  409. int rootL = roots[versionL];
  410. int rootR = roots[versionR];
  411. return queryKthRec(rootL, rootR, 1, N, k);
  412. }
  413. };
  414.  
  415. // ===================================================================
  416. // EXAMPLE USAGE
  417. // ===================================================================
  418.  
  419. int main() {
  420. ios::sync_with_stdio(false);
  421. cin.tie(nullptr);
  422.  
  423. // ---------------------------------------------------------------
  424. // Demo 1: Persistent tree with range add and sum
  425. // ---------------------------------------------------------------
  426. vector<ll> arr = {1, 2, 3, 4, 5};
  427. PersistentSegTreeLazy pst(arr);
  428.  
  429. int root0 = pst.getRoot(0);
  430. cout << "Initial sum [1,3] = " << pst.queryRangeSum(root0, 1, 3) << "\n"; // 9
  431.  
  432. int root1 = pst.updateRangeAdd(root0, 1, 3, 10); // add 10 to indices 1..3
  433. pst.addVersion(root1);
  434. cout << "After add, sum [1,3] = " << pst.queryRangeSum(root1, 1, 3) << "\n"; // 39
  435.  
  436. // Find the first index where prefix sum > 25
  437. int idx = pst.findFirstPrefixGreater(root1, 25);
  438. cout << "First index with prefix > 25 = " << idx << "\n"; // 2 (prefix up to 2 = 26)
  439.  
  440. // ---------------------------------------------------------------
  441. // Demo 2: Persistent tree with range set (assignment)
  442. // ---------------------------------------------------------------
  443. PersistentSegTreeSet pstSet(arr);
  444. int r0 = pstSet.getRoot(0);
  445. cout << "Initial sum [0,4] = " << pstSet.queryRangeSum(r0, 0, 4) << "\n"; // 15
  446.  
  447. int r1 = pstSet.updateRangeSet(r0, 0, 2, 100); // set first three to 100
  448. pstSet.addVersion(r1);
  449. cout << "After setting [0,2] to 100, sum all = " << pstSet.queryRangeSum(r1, 0, 4) << "\n"; // 100+100+100+4+5 = 309
  450.  
  451. // ---------------------------------------------------------------
  452. // Demo 3: K‑th smallest using the count tree
  453. // ---------------------------------------------------------------
  454. vector<int> values = {5, 2, 8, 2, 9};
  455. // Coordinate compression
  456. vector<int> sortedVals = values;
  457. sort(sortedVals.begin(), sortedVals.end());
  458. sortedVals.erase(unique(sortedVals.begin(), sortedVals.end()), sortedVals.end());
  459.  
  460. PersistentSegTreeCount pstCount((int)sortedVals.size());
  461. int emptyRoot = pstCount.getRoot(0);
  462. vector<int> countRoots = {emptyRoot}; // root[0] = empty prefix
  463.  
  464. for (int x : values) {
  465. int pos = lower_bound(sortedVals.begin(), sortedVals.end(), x) - sortedVals.begin() + 1; // 1‑indexed
  466. int newRoot = pstCount.updatePoint(countRoots.back(), pos, 1);
  467. countRoots.push_back(newRoot);
  468. }
  469.  
  470. // Query 2nd smallest in subarray [1, 3] (0‑based) -> values {2,8,2} -> sorted {2,2,8}
  471. // We need versions: L = 1 (prefix before index 1), R = 4 (prefix up to index 3).
  472. int kthIdx = pstCount.queryKth(1, 4, 2);
  473. cout << "2nd smallest in [1,3] = " << sortedVals[kthIdx - 1] << "\n"; // should be 2
  474.  
  475. return 0;
  476. }
Success #stdin #stdout 0.01s 5320KB
stdin
Standard input is empty
stdout
Initial sum [1,3] = 9
After add, sum [1,3] = 39
First index with prefix > 25 = 2
Initial sum [0,4] = 15
After setting [0,2] to 100, sum all = 309
2nd smallest in [1,3] = 9