来个简单好写做法。
设 \(g(v)\) 表示最少可以用多少个长度为 \(d\) 的区间覆盖所有 \(\geq v\) 的位置,那么题目转化为求 \(\sum_{v\geq 1}g(v)\)。
求单个 \(g(v)\) 是简单的,可以直接贪心,每次选最靠左的未被覆盖的点作为新区间的左端点即可。
贪心不好刻画,考虑直接 DP。令 \(f_{i,v}\) 表示考虑 \(a[i..n]\),最少可以用多少个长度为 \(d\) 的区间覆盖所有 \(\geq v\) 的位置。转移是容易的:
不妨设 \(a\) 排序并去重后得到 \(v_1,\cdots,v_m\),那么对于 \(v\in(v_{j-1},v_j]\),\(f_{i,v}\) 的取值是一样的,不妨压缩状态,将其表示成 \(f_{i,j}\)。那么设 \(a_i=v_k\),转移变为
考虑把 \(f_i\) 看作长度为 \(m\) 的序列,那么我们相当于取 \(f_{i+d}[1..k]\) 整体 \(+1\) 再接上 \(f_{i+1}[k+1..m]\),得到 \(f_i\)。
考虑用可持久化线段树维护。节点上维护权值和 \(\sum_{j=l}^r(v_j-v_{j-1})f_{i,j}\)。递归到 \([l,r]\) 时,若 \(r\leq k\),就把 \(f_{i+d}\) 的对应节点整体 \(+1\) 后复用;若 \(l>k\) 就拿 \(f_{i+1}\) 的对应节点复用;否则新建节点,向两边递归。显然每层至多有一个节点跨过分界线,因此每次只会新建 \(\mathcal{O}(\log{m})\) 个节点。
对于整体 \(+1\),我们不妨在每个节点上维护 \((p,tag)\),表示这个节点复用的是节点 \(p\),且整体加 \(tag\)。整体加 \(tag\) 后权值和会增加 \(tag(v_r-v_{l-1})\)。
视 \(n,m\) 同阶,时空复杂度均为 \(\mathcal{O}(n\log{n})\)。
代码很好写。
代码
#include <bits/stdc++.h>using namespace std;using ll = long long;
using i128 = __int128;
using ui = unsigned int;
using ull = unsigned long long;
using u128 = unsigned __int128;
using ld = long double;
using pii = pair<int, int>;
const int MAXN = 5e5 + 5;template<typename T> T lowbit(T x) { return x & -x; }
template<typename T> void chkMin(T &x, T y) { x = y < x ? y : x; }
template<typename T> void chkMax(T &x, T y) { x = x < y ? y : x; }
constexpr int lg2(ll x) { return 63 ^ __builtin_clzll(x); }
constexpr ll bitCeil(ll x) { return x == 1 ? 1ll : 1ll << lg2(x - 1) + 1; }int tc, n, d, m, a[MAXN];
pii rt[MAXN];
vector<int> disc;struct SegTree {static const int MAXC = 1e7 + 5;int tot;pii ls[MAXC], rs[MAXC];ll sum[MAXC];ll calc(pii p, int l, int r) {return sum[p.first] + (ll)p.second * (disc[r] - disc[l - 1]);}pii solve(pii p, pii q, int l, int r, int x) {if (r <= x) {++p.second;return p;}if (l > x) return q;int mid = l + r >> 1;auto [pid, ptg] = p;auto [qid, qtg] = q;pii pl = {ls[pid].first, ls[pid].second + ptg};pii ql = {ls[qid].first, ls[qid].second + qtg};pii L = solve(pl, ql, l, mid, x);pii pr = {rs[pid].first, rs[pid].second + ptg};pii qr = {rs[qid].first, rs[qid].second + qtg};pii R = solve(pr, qr, mid + 1, r, x);int cur = ++tot;ls[cur] = L;rs[cur] = R;sum[cur] = calc(L, l, mid) + calc(R, mid + 1, r);return {cur, 0};}
} sgt;int main() {ios::sync_with_stdio(false);cin.tie(nullptr);cin >> tc;while (tc--) {cin >> n >> d;for (int i = 1; i <= n; ++i) cin >> a[i];disc = vector<int>(a + 1, a + n + 1);disc.emplace_back(0);sort(disc.begin(), disc.end());disc.erase(unique(disc.begin(), disc.end()), disc.end());m = disc.size() - 1;rt[n + 1] = {0, 0};sgt.tot = 0;for (int i = n; i; --i) {if (!a[i]) {rt[i] = rt[i + 1];continue;}int v = lower_bound(disc.begin(), disc.end(), a[i]) - disc.begin();rt[i] = sgt.solve(rt[min(i + d, n + 1)], rt[i + 1], 1, m, v);}cout << sgt.calc(rt[1], 1, m) << '\n';}return 0;
}