Skip to Content

Cake Circle

Explicación

Primero, dado algún conjunto de pasteles SS, siempre es óptimo disponerlos en orden ordenado por color. Para demostrarlo, consideremos agregar inductivamente los pasteles en orden creciente. Se puede mostrar que colocar el pastel entre el mayor agregado hasta ahora y el menor siempre es óptimo. Así, el color total es 2(max(S)min(S))2(\max(S) - \min(S)).

Ordenamos los pasteles por color. Para cualquier conjunto de pasteles, el más pequeño es el de más a la izquierda y el más grande es el de más a la derecha. La tarea ahora es elegir algún rango [l,r][l, r] que maximice la suma de los mm mayores valores de belleza, menos 2(CrCl)2 \cdot (C_r - C_l).

Sea opt(l)\texttt{opt}(l) el extremo derecho óptimo para algún extremo izquierdo. Podemos observar que opt\texttt{opt} es una función no decreciente (la demostración se deja como ejercicio para el lector). Esto motiva una solución de divide y vencerás.

Supongamos que sabemos que opt(i)[p,q]\texttt{opt}(i) \in [p, q] para i[l,r]i \in [l, r]. Si p=qp = q, entonces hemos terminado. En caso contrario, sea m=l+r2m = \lfloor \frac{l+r}{2}\rfloor, y calculemos opt(m)\texttt{opt}(m). Ahora sabemos que opt(i)\texttt{opt}(i) está en el rango [p,opt(m)][p, \texttt{opt(m)}] para i[l,m)i \in [l, m), y en el rango [opt(m),q][\texttt{opt}(m), q] para i(m,r]i \in (m, r]. Luego podemos recurrir. Por el teorema maestro , esto corre en tiempo O~(n)\tilde{\mathcal{O}}(n).

El único paso que queda es calcular la suma de los mm mayores valores en un rango. Esto se puede hacer con varias estructuras de datos en tiempo logarítmico.

Complejidad temporal: O(nlog2n)\mathcal{O}(n\log^2n)

Implementación

#include <bits/stdc++.h> using namespace std; namespace std { template <class Fun> class y_combinator_result { Fun fun_; public: template <class T> explicit y_combinator_result(T &&fun) : fun_(std::forward<T>(fun)) {} template <class... Args> decltype(auto) operator()(Args &&...args) { return fun_(std::ref(*this), std::forward<Args>(args)...); } }; template <class Fun> decltype(auto) y_combinator(Fun &&fun) { return y_combinator_result<std::decay_t<Fun>>(std::forward<Fun>(fun)); } } // namespace std class WaveletTree { int n; vector<long long> vmap; vector<vector<int>> tree; vector<vector<long long>> psum; template <typename I> void build(I begin, I end, int t, int vl, int vr) { if (vl == vr) { int nn = end - begin + 1; psum[t].reserve(nn); psum[t].push_back(0); for (auto it = begin; it != end; it++) psum[t].push_back(psum[t].back() + vmap[*it]); return; } int vm = (vl + vr) / 2, nn = end - begin + 1; tree[t].reserve(nn); psum[t].reserve(nn); tree[t].push_back(0); psum[t].push_back(0); for (auto it = begin; it != end; it++) { tree[t].push_back(tree[t].back() + (*it <= vm)); psum[t].push_back(psum[t].back() + vmap[*it]); } auto pivot = stable_partition(begin, end, [vm](int x) { return x <= vm; }); build(begin, pivot, t * 2, vl, vm); build(pivot, end, t * 2 + 1, vm + 1, vr); } long long query(int l, int r, int k, int t, int vl, int vr) { if (vl == vr) return k * vmap[vl]; int vm = (vl + vr) / 2, lv = tree[t][r] - tree[t][l - 1]; if (k <= lv) { return query(tree[t][l - 1] + 1, tree[t][r], k, t * 2, vl, vm); } else { return psum[t * 2][tree[t][r]] - psum[t * 2][tree[t][l - 1]] + query(l - tree[t][l - 1], r - tree[t][r], k - lv, t * 2 + 1, vm + 1, vr); } } public: WaveletTree(vector<long long> vec) { init(vec.begin(), vec.end()); } // resets data structure to initial state template <typename I> void init(I begin, I end) { map<long long, int> m; for (auto it = begin; it != end; it++) m[*it] = 0; n = 0; vmap.resize(m.size()); for (auto &[k, v] : m) vmap[v = n++] = k; for (auto it = begin; it != end; it++) *it = m[*it]; tree.resize(4 * n); psum.resize(4 * n); build(begin, end, 1, 0, n); } // returns sum of k smallest elements in range [l, r] long long query(int l, int r, int k) { assert(k <= r - l + 1); return query(l + 1, r + 1, k, 1, 0, n); } }; int main() { int N, M; cin >> N >> M; vector<pair<long long, long long>> A(N); for (auto &[v, c] : A) { cin >> v >> c; v = -v; } sort(A.begin(), A.end(), [](auto a, auto b) { return a.second < b.second; }); vector<long long> V(N), C(N); for (int i = 0; i < N; i++) { tie(V[i], C[i]) = A[i]; } WaveletTree wt(V); vector<int> rpos(N - M + 1); y_combinator([&](auto self, int pl, int pr, int ql, int qr) -> void { if (pr < pl) { return; } if (ql == qr) { for (int i = pl; i <= pr; i++) rpos[i] = ql; return; } int pm = (pl + pr) / 2, qm = -1; long long val = LLONG_MIN; for (int i = ql; i <= qr; i++) { if (i - pm + 1 < M) { continue; } long long cur = -wt.query(pm, i, M) - 2 * (C[i] - C[pm]); if (val < cur) { val = cur; qm = i; } } rpos[pm] = qm; self(pl, pm - 1, ql, qm); self(pm + 1, pr, qm, qr); })(0, N - M, M - 1, N - 1); long long ans = LLONG_MIN; for (int l = 0; l < N - M + 1; l++) { int r = rpos[l]; ans = max(ans, -wt.query(l, r, M) - 2 * (C[r] - C[l])); } cout << ans << endl; }