Skip to Content

Más aplicaciones del Árbol de Segmentos

Recursos
FuenteRecursoNotas
CF EDUSegment Tree Pt 1 Step 2

ambos temas

cp-algoSegment Tree - More complex queries

Incluye estas dos aplicaciones y más.

Caminar sobre un Árbol de Segmentos

HechoFuenteNombreDificultadTagsSolución
CSESHotel QueriesFácilPURQen el módulo

Queremos soportar consultas de la siguiente forma sobre un arreglo a1,,aNa_1,\ldots,a_N (junto con actualizaciones puntuales).

Hallar el primer ii tal que aixa_i\ge x.

Por supuesto, esto se puede hacer en tiempo O(log2N)\mathcal{O}(\log^2N) con un árbol de segmentos de máximos y búsqueda binaria sobre el primer ii tal que max(a1,,ai)x\max(a_1,\ldots,a_i)\ge x. Pero hay que intentar hacerlo en tiempo O(logN)\mathcal{O}(\log N).

Solución - Hotel Queries

En vez de hacer búsqueda binaria y consultar el árbol de segmentos por separado, ¡hagámoslo junto!

Supongamos que sabemos que la respuesta está en algún rango [l,r][l, r]. Sea mid=l+r2mid = \left \lfloor{\frac{l + r}{2}}\right \rfloor.

Si max(al,,amid)x\max(a_l, \dots, a_{mid}) \geq x, entonces sabemos que la respuesta está en el rango [l,mid][l, mid]. Si no, la respuesta está en el rango [mid+1,r][mid + 1, r].

Imaginemos que el árbol de segmentos es un árbol de decisión. Empezamos en la raíz y bajamos. Cuando estamos en algún nodo que contiene max(al,,ar)\max(a_l, \dots, a_r) y sabemos que la respuesta está en el rango [l,r][l, r], nos movemos al hijo izquierdo si max(al,,amid)x\max(a_l, \dots, a_{mid}) \geq x; si no, nos movemos al hijo derecho.

Esto es conveniente porque max(al,,amid)\max(a_l, \dots, a_{mid}) ya está guardado en el hijo izquierdo, así que lo podemos hallar en tiempo O(1)\mathcal{O}(1).

Complejidad temporal: O(N+QlogN)\mathcal{O}(N + Q\log{N})

#include <bits/stdc++.h> using namespace std; const int MAXN = 200001; int n; int segtree[4 * MAXN], a[MAXN]; void build(int l = 1, int r = n, int node = 1) { if (l == r) segtree[node] = a[l]; else { int mid = (l + r) / 2; build(l, mid, node * 2); build(mid + 1, r, node * 2 + 1); segtree[node] = max(segtree[node * 2], segtree[node * 2 + 1]); } } void queryupdate(int val, int l = 1, int r = n, int node = 1) { if (l == r) { segtree[node] -= val; cout << l << ' '; } else { int mid = (l + r) / 2; if (segtree[node * 2] >= val) queryupdate(val, l, mid, node * 2); else queryupdate(val, mid + 1, r, node * 2 + 1); segtree[node] = max(segtree[node * 2], segtree[node * 2 + 1]); } } int main() { iostream::sync_with_stdio(false); cin.tie(0); int q; cin >> n >> q; for (int i = 1; i <= n; i++) cin >> a[i]; build(); while (q--) { int x; cin >> x; if (segtree[1] < x) cout << "0 "; else queryupdate(x); } return 0; }
import java.io.*; import java.util.*; public class Main { static int n; static int[] segtree; static int[] a; static void build(int node, int l, int r) { if (l == r) { segtree[node] = a[l]; return; } int mid = (l + r) / 2; build(node * 2, l, mid); build(node * 2 + 1, mid + 1, r); segtree[node] = Math.max(segtree[node * 2], segtree[node * 2 + 1]); } static void queryUpdate(int node, int l, int r, int val, StringBuilder out) { if (l == r) { segtree[node] -= val; out.append(l).append(' '); return; } int mid = (l + r) / 2; if (segtree[node * 2] >= val) { queryUpdate(node * 2, l, mid, val, out); } else { queryUpdate(node * 2 + 1, mid + 1, r, val, out); } segtree[node] = Math.max(segtree[node * 2], segtree[node * 2 + 1]); } public static void main(String[] args) throws Exception { Kattio io = new Kattio(); n = io.nextInt(); int q = io.nextInt(); a = new int[n + 1]; for (int i = 1; i <= n; i++) a[i] = io.nextInt(); segtree = new int[4 * n]; build(1, 1, n); StringBuilder out = new StringBuilder(); while (q-- > 0) { int x = io.nextInt(); if (segtree[1] < x) { out.append("0 "); } else { queryUpdate(1, 1, n, x, out); } } io.println(out.toString()); io.close(); } // CodeSnip{Kattio} }
HechoFuenteNombreDificultadTagsSolución
SPOJOrdered SetFácilPURSen el módulo

Solución - Ordered Set

Primero, comprimimos coordenadas de todos los valores y llevamos la frecuencia de cada valor en el arreglo sobre el que construimos el árbol de segmentos. Podemos responder las consultas INSERT\texttt{INSERT}, DELETE\texttt{DELETE} y COUNT\texttt{COUNT} usando el árbol de segmentos de sumas del módulo PURS. Responder las consultas K-TH\texttt{K-TH} en tiempo O(logN)\mathcal{O}(\log{N}) requiere caminar sobre el árbol de segmentos.

Sea freq[i]\texttt{freq}[i] la cantidad de veces que ii aparece en nuestro conjunto. En nuestro arreglo de valores con coordenadas comprimidas, queremos hallar el primer índice xx tal que i=0xfreq[i]\sum_{i=0}^{x} \texttt{freq}[i] es K\geq K.

En el problema anterior, caminábamos sobre el máximo de prefijos de nuestro arreglo. Acá, caminamos sobre la suma de prefijos de nuestro arreglo. La diferencia es que si sabemos que la respuesta está en [l,r][l, r], entonces chequear el máximo en los hijos izquierdo y derecho alcanza para saber si hay que caminar a la izquierda o a la derecha. Pero con la suma, también hay que considerar el rango [1,l)[1, l) y cómo puede afectar la respuesta.

Para llevar el registro del prefijo [1,l)[1, l), primero ponemos el resultado de prefijo en un valor neutro. Por valor neutro entendemos la identidad de la operación (para la suma es 0, para la multiplicación es 1, etc.). Cada vez que caminamos a la derecha, sumamos el valor del hijo izquierdo a nuestro resultado de prefijo.

Complejidad temporal: O(QlogQ)\mathcal{O}(Q\log{Q})

#include <bits/stdc++.h> using namespace std; template <class T> class SumSegmentTree { private: const T DEFAULT = 0; int len; vector<T> segtree; T combine(const T &a, const T &b) { return a + b; } void build(const vector<T> &arr, int at, int at_left, int at_right) { if (at_left == at_right) { segtree[at] = arr[at_left]; return; } int mid = (at_left + at_right) / 2; build(arr, 2 * at, at_left, mid); build(arr, 2 * at + 1, mid + 1, at_right); segtree[at] = combine(segtree[2 * at], segtree[2 * at + 1]); } void set(int ind, T val, int at, int at_left, int at_right) { if (at_left == at_right) { segtree[at] = val; return; } int mid = (at_left + at_right) / 2; if (ind <= mid) { set(ind, val, 2 * at, at_left, mid); } else { set(ind, val, 2 * at + 1, mid + 1, at_right); } segtree[at] = combine(segtree[2 * at], segtree[2 * at + 1]); } T range_sum(int start, int end, int at, int at_left, int at_right) { if (at_right < start || end < at_left) { return DEFAULT; } if (start <= at_left && at_right <= end) { return segtree[at]; } int mid = (at_left + at_right) / 2; T left_res = range_sum(start, end, 2 * at, at_left, mid); T right_res = range_sum(start, end, 2 * at + 1, mid + 1, at_right); return combine(left_res, right_res); } int walk(int at, int at_left, int at_right, int max_val, int pref_res) { if (at_left == at_right) { return at_left; } int mid = (at_left + at_right) / 2; int sum_left = pref_res + segtree[2 * at]; if (sum_left >= max_val) { return walk(2 * at, at_left, mid, max_val, pref_res); } return walk(2 * at + 1, mid + 1, at_right, max_val, sum_left); } public: SumSegmentTree(int len) : len(len) { segtree = vector<T>(len * 4, DEFAULT); }; SumSegmentTree(const vector<T> &arr) : len(arr.size()) { segtree = vector<T>(len * 4, DEFAULT); build(arr, 1, 0, len - 1); } void set(int ind, T val) { set(ind, val, 1, 0, len - 1); } T range_sum(int start, int end) { return range_sum(start, end, 1, 0, len - 1); } /** @return primer i tal que la suma de prefijos hasta i >= val */ int walk(int val) { return walk(1, 0, len - 1, val, 0); } }; int main() { int query_num; cin >> query_num; vector<int> indices; vector<pair<char, int>> updates(query_num); for (int i = 0; i < query_num; i++) { char type; int val; cin >> type >> val; updates[i] = {type, val}; indices.push_back(val); } sort(begin(indices), end(indices)); indices.erase(unique(begin(indices), end(indices)), end(indices)); /** @return ubicación con coordenadas comprimidas de v */ const auto id = [&](int v) -> int { return lower_bound(begin(indices), end(indices), v) - begin(indices); }; vector<int> freq(indices.size()); SumSegmentTree<int> st(indices.size()); for (const auto [type, val] : updates) { int loc = id(val); if (type == 'I') { if (freq[loc] == 0) { freq[loc] = 1; st.set(loc, 1); } } else if (type == 'D') { if (freq[loc] == 1) { freq[loc] = 0; st.set(loc, 0); } } else if (type == 'K') { if (st.range_sum(0, indices.size() - 1) < val) { cout << "invalid" << '\n'; } else { cout << st.walk(val) << '\n'; } } else if (type == 'C') { cout << st.range_sum(0, loc - 1) << '\n'; } } }
import java.io.*; import java.util.*; public class Main { static class SumSegTree { int n; int[] tree; SumSegTree(int n) { this.n = n; tree = new int[4 * n]; } void set(int idx, int val, int node, int l, int r) { if (l == r) { tree[node] = val; return; } int mid = (l + r) / 2; if (idx <= mid) set(idx, val, node * 2, l, mid); else set(idx, val, node * 2 + 1, mid + 1, r); tree[node] = tree[node * 2] + tree[node * 2 + 1]; } void set(int idx, int val) { set(idx, val, 1, 0, n - 1); } int rangeSum(int ql, int qr, int node, int l, int r) { if (qr < l || r < ql) return 0; if (ql <= l && r <= qr) return tree[node]; int mid = (l + r) / 2; return rangeSum(ql, qr, node * 2, l, mid) + rangeSum(ql, qr, node * 2 + 1, mid + 1, r); } int rangeSum(int l, int r) { if (l > r) return 0; return rangeSum(l, r, 1, 0, n - 1); } int walk(int node, int l, int r, int k, int pref) { if (l == r) return l; int mid = (l + r) / 2; int leftSum = pref + tree[node * 2]; if (leftSum >= k) { return walk(node * 2, l, mid, k, pref); } return walk(node * 2 + 1, mid + 1, r, k, leftSum); } int walk(int k) { return walk(1, 0, n - 1, k, 0); } } public static void main(String[] args) throws Exception { Kattio io = new Kattio(); int q = io.nextInt(); char[] type = new char[q]; int[] val = new int[q]; ArrayList<Integer> coords = new ArrayList<>(); for (int i = 0; i < q; i++) { type[i] = io.next().charAt(0); val[i] = io.nextInt(); coords.add(val[i]); } Collections.sort(coords); coords = new ArrayList<>(new LinkedHashSet<>(coords)); int m = coords.size(); HashMap<Integer, Integer> id = new HashMap<>(); for (int i = 0; i < m; i++) id.put(coords.get(i), i); int[] freq = new int[m]; SumSegTree st = new SumSegTree(m); for (int i = 0; i < q; i++) { int loc = id.get(val[i]); if (type[i] == 'I') { if (freq[loc] == 0) { freq[loc] = 1; st.set(loc, 1); } } else if (type[i] == 'D') { if (freq[loc] == 1) { freq[loc] = 0; st.set(loc, 0); } } else if (type[i] == 'K') { if (st.rangeSum(0, m - 1) < val[i]) { io.println("invalid"); } else { io.println(st.walk(val[i])); } } else { io.println(st.rangeSum(0, loc - 1)); } } io.close(); } // CodeSnip{Kattio} }

Problemas

HechoFuenteNombreDificultadTagsSolución
Old GoldSeatingNormalSolución
CFSerge and Dining RoomNormal

Funciones combinadoras no conmutativas

Antes solo consideramos operaciones conmutativas como + y max. Sin embargo, los árboles de segmentos permiten responder consultas de rango para cualquier operación asociativa.

HechoFuenteNombreDificultadTagsSolución
YSPoint Set Range CompositeFácilPURQen el módulo
HechoFuenteNombreDificultadTagsSolución
CSESSubarray Sum QueriesNormalPURQen el módulo

Solución - Point Set Range Composite

El árbol de segmentos del módulo prerrequisito debería alcanzar. También se pueden usar dos BIT como se describe acá , aunque es más complicado.

using T = AR<mi, 2>; T comb(const T &a, const T &b) { return {a[0] * b[0], a[1] * b[0] + b[1]}; } template <class T> struct BIT { const T ID = {1, 0}; int SZ = 1; V<T> x, bit[2]; void init(int N) { while (SZ <= N) SZ *= 2; x = V<T>(SZ + 1, ID); F0R(i, 2) bit[i] = x; FOR(i, 1, N + 1) re(x[i]); build(); } void build() { FOR(i, 1, SZ) { bit[0][i] = comb(bit[0][i], x[i]); int j = i + (i & -i); assert(j <= SZ); bit[0][j] = comb(bit[0][j], bit[0][i]); } ROF(i, 1, SZ) { bit[1][i] = comb(x[i], bit[1][i]); int j = i - (i & -i); bit[1][j] = comb(bit[1][i], bit[1][j]); } } void upd0(int p) { T lans = ID, rans = ID; for (int P = p, lo = p - 1, hi = p + 1; P < SZ; P += P & -P) { for (; hi < P; hi += hi & -hi) rans = comb(rans, bit[1][hi]); for (; lo > P - (P & -P); lo -= lo & -lo) lans = comb(bit[0][lo], lans); assert(lo == P - (P & -P)); bit[0][P] = comb(lans, x[p]); if (p != P) bit[0][P] = comb(bit[0][P], comb(rans, x[P])); } } void upd1(int p) { T lans = ID, rans = ID; for (int P = p, lo = p - 1, hi = p + 1; P; P -= P & -P) { for (; hi < P + (P & -P); hi += hi & -hi) rans = comb(rans, bit[1][hi]); for (; lo > P; lo -= lo & -lo) lans = comb(bit[0][lo], lans); assert(hi == P + (P & -P)); bit[1][P] = comb(x[p], rans); if (p != P) bit[1][P] = comb(comb(x[P], lans), bit[1][P]); } } void upd(int p, T u) { assert(0 < p && p < SZ); x[p] = u; upd0(p); upd1(p); } T query(int a, int b) { assert(0 < a && a <= b && b < SZ); T lans = ID, rans = ID; for (int A; (A = a + (a & -a)) <= b; a = A) lans = comb(lans, bit[1][a]); for (int B; (B = b - (b & -b)) >= a; b = B) rans = comb(bit[0][b], rans); assert(a == b); return comb(comb(lans, x[a]), rans); } }; BIT<T> B; int N, Q; int main() { setIO(); re(N, Q); B.init(N); F0R(i, Q) { int t, p, c, d; re(t, p, c, d); ++p; if (t == 0) { B.upd(p, {c, d}); } else { T res = B.query(p, c); ps(res[0] * d + res[1]); } } }

Solución - Subarray Sum Queries

Pista: en cada nodo del árbol de segmentos hay que guardar cuatro piezas de información.

En cada nodo del árbol de segmentos que guarda información sobre el rango [l,r][l, r] almacenamos lo siguiente:

  • La suma máxima de subarreglo en el rango [l,r][l, r]. (Sea esto GG)
  • La suma máxima de subarreglo en el rango [l,r][l, r] si debe contener ala_l. (Sea esto LL)
  • La suma máxima de subarreglo en el rango [l,r][l, r] si debe contener ara_r. (Sea esto RR)
  • La suma total del rango. (Sea esto SS)

Cuando combinamos dos nodos uu (hijo izquierdo) y vv (hijo derecho) para formar el nodo ww,

  • w.G=max(u.G,v.G,u.R+v.L)w.G = \max(u.G, v.G, u.R + v.L)
  • w.L=max(u.L,u.S+v.L)w.L = \max(u.L, u.S + v.L)
  • w.R=max(u.R+v.S,v.R)w.R = \max(u.R + v.S, v.R)
  • w.S=u.S+v.Sw.S = u.S + v.S

Así podemos manejar actualizaciones y consultas de forma eficiente.

#include <bits/stdc++.h> typedef long long ll; using namespace std; const ll MAXN = 200001; struct Node { ll g_max, l_max, r_max, sum; Node operator+(Node b) { return {max(max(g_max, b.g_max), r_max + b.l_max), max(l_max, sum + b.l_max), max(b.r_max, r_max + b.sum), sum + b.sum}; } }; ll n, a[MAXN]; Node segtree[4 * MAXN]; void build(ll l = 1, ll r = n, ll node = 1) { if (l == r) segtree[node] = {max(0ll, a[l]), max(0ll, a[l]), max(0ll, a[l]), a[l]}; else { ll mid = (l + r) / 2; build(l, mid, node * 2); build(mid + 1, r, node * 2 + 1); segtree[node] = segtree[node * 2] + segtree[node * 2 + 1]; } } void update(ll pos, ll val, ll l = 1, ll r = n, ll node = 1) { if (l == r) segtree[node] = {max(0ll, val), max(0ll, val), max(0ll, val), val}; else { ll mid = (l + r) / 2; if (pos > mid) update(pos, val, mid + 1, r, node * 2 + 1); else update(pos, val, l, mid, node * 2); segtree[node] = segtree[node * 2] + segtree[node * 2 + 1]; } } int main() { iostream::sync_with_stdio(false); cin.tie(0); ll q; cin >> n >> q; for (int i = 1; i <= n; i++) cin >> a[i]; build(); while (q--) { ll x, y; cin >> x >> y; update(x, y); cout << segtree[1].g_max << '\n'; } return 0; }

Problemas

HechoFuenteNombreDificultadTagsSolución
CSESPizzeria QueriesFácilSolución
Old GoldMarathonFácil
CSESBit InversionsFácilPURQSolución
CF01 (Hard Version)FácilPURQ
PlatinumHigh Card Low CardFácilPURQ, GreedySolución
Old GoldOptimal MilkingNormal
POI2014 - CardsNormalSolución
PlatinumPareidoliaNormalMatrix, PURQSolución
COCI2021 - SjeckanjeDifícilPURQSolución
Balkan OI2018 - ElectionDifícilSolución