题解:P15947 [JOI Final 2026] 集邮 5 / Collecting Stamps 5
一个非常点分治的好题。
题目就是说,让你求对于每个
更公式化地说,你需要求有多少条路径
考虑把不等式转化一下:
既然转化成了含有最近公共祖先的式子,就不难考虑到点分治。
将式子拆开分成两个部分,一个部分是
两个部分的式子分别是
这下就很明显了,对于一个分治中心
那如何判断呢?
简单地,其充要条件就是一条路径的最小值是否满足条件就行了。
证明是简单的。充分的证明是显然的,你一条路径中值最小的点满足
然后你通过总体满足条件的数量减去所在子树满足条件的数量的写法,发现只得到了
加上了
然后就做完了。
提前说一下,这个过的最大时间为
::::success[代码]
提交记录
#include <bits/stdc++.h>
using namespace std;
const int N = 5e5+10;
const int M = 1e6+10;
const int del = 5e5;
const int V = 1e6;
const int inf = 0x3f3f3f3f;
struct node {
int ls, rs, sum;
};
struct edge {
int u, from, now;
};
class SegTree {
public:
int rt[V], las[V];
node seg[35*V];
int base, nowpos;
public:
void Copy(int x, int y) {
seg[x] = seg[y];
}
int doBuild(int k, int l, int r) {
base = max(base, k);
if (l == r) return k;
int mid = (l + r) >> 1;
seg[k].ls = doBuild(k*2, l, mid);
seg[k].rs = doBuild(k*2+1, mid+1, r);
return k;
}
int doChange(int k, int l, int r, int x, int dx) {
int now = ++nowpos; Copy(now, k);
if (l == r) {
seg[now].sum += dx;
return now;
}
int mid = (l + r) >> 1;
if (x <= mid) seg[now].ls = doChange(seg[now].ls, l, mid, x, dx);
else seg[now].rs = doChange(seg[now].rs, mid+1, r, x, dx);
seg[now].sum = seg[seg[now].ls].sum + seg[seg[now].rs].sum;
return now;
}
int doQuery(int k, int l, int r, int x, int y) {
if (r < x || y < l) return 0;
if (x <= l && r <= y) return seg[k].sum;
int mid = (l + r) >> 1;
return doQuery(seg[k].ls, l, mid, x, y) + doQuery(seg[k].rs, mid+1, r, x, y);
}
} global, part;
int n, d;
int t[N];
int cnt, head[N];
int nxt[M], to[M];
int dep[N];
int now;
int siz[N], wight[N];
int rikka[N];
int nashi[N];
vector<int> vec;
bool vis[N];
int output[N];
vector<pair<int, int> > all;
queue<edge> q;
__inline void read(int &x) {
x = 0; int f = 1;
char ch = getchar_unlocked();
while (!(ch >= '0' && ch <= '9')) {
if (ch == '-') f = -1;
ch = getchar_unlocked();
}
while (ch >= '0' && ch <= '9') {
x = x * 10 + (ch - '0');
ch = getchar_unlocked();
} x *= f;
}
__inline void connect(int u, int v) {
cnt++, nxt[cnt] = head[u];
head[u] = cnt, to[cnt] = v;
}
void getCentroid(int u, int from, int &ctre) {
siz[u] = 1, wight[u] = 1;
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (v == from || vis[v])
continue;
getCentroid(v, u, ctre);
wight[u] = max(wight[u], siz[v]);
siz[u] += siz[v];
}
wight[u] = max(wight[u], now - siz[u]);
if (wight[u] <= now / 2) ctre = u;
}
void dfs2(int u, int from, int now) {
dep[u] = dep[from] + 1;
now = min(now, t[u] - dep[u]);
rikka[u] = now; siz[u] = 1;
all.push_back({dep[u], u});
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (v == from || vis[v])
continue;
dfs2(v, u, now);
siz[u] += siz[v];
}
}
__inline void bfs(int u, int from, int now) {
q.push({u, from, now});
while (!q.empty()) {
int u = q.front().u, from = q.front().from;
int now = min(q.front().now, t[u] + dep[u]);
nashi[u] = now; vec.emplace_back(u); q.pop();
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (v == from || vis[v]) continue;
q.push({v, u, now});
}
}
}
void calc(int u, int from, int num) {
int centre = 0; now = num;
getCentroid(u, from, centre);
u = centre; dep[0] = -1, dfs2(u, 0, inf);
sort(all.begin(), all.end());
int Max = 0;
global.nowpos = global.base; int cur = 0;
for (pair<int, int> kv : all) {
global.las[kv.first] = ++cur; Max = max(Max, kv.first);
global.rt[cur] = global.doChange(global.rt[cur-1], 0, V, del+rikka[kv.second], 1);
}
int y = global.las[d];
if (y == -1) y = cur;
output[u] += global.doQuery(global.rt[y], 0, V, 0, max(0, del-dep[u]));
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (vis[v]) continue;
bfs(v, u, t[u]+dep[u]);
int cc = 0, Max = 0;
part.nowpos = part.base;
for (int nd : vec) {
part.las[dep[nd]] = ++cc; Max = dep[nd];
part.rt[cc] = part.doChange(part.rt[cc-1], 0, V, del+rikka[nd], 1);
}
for (int nd : vec) {
if (d - dep[nd] >= 0) {
int x = part.las[d-dep[nd]], y = global.las[d-dep[nd]];
if (x == -1) x = cc;
if (y == -1) y = cur;
if (nashi[nd] <= dep[nd])
output[nd] += global.doQuery(global.rt[y], 0, V, 0, V) - part.doQuery(part.rt[x], 0, V, 0, V);
else
output[nd] += global.doQuery(global.rt[y], 0, V, 0, max(0, dep[nd]-2*dep[u]+del)) -
part.doQuery(part.rt[x], 0, V, 0, max(0, dep[nd]-2*dep[u]+del));
}
}
for (int i = 1;i<= Max;i++) part.las[i] = -1;
part.las[0] = 0;
vec.clear();
}
for (int i = 1;i<= Max;i++) {
global.las[i] = -1;
} global.las[0] = 0;
all.clear();
vis[u] = 1;
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (vis[v]) continue;
calc(v, u, siz[v]);
}
}
signed main() {
read(n), read(d);
for (int i = 1;i<= n;i++) {
read(t[i]);
} int u, v;
for (int i = 1;i< n;i++) {
read(u), read(v);
connect(u, v);
connect(v, u);
}
memset(part.las, -1, sizeof part.las);
part.las[0] = 0;
memset(global.las, -1, sizeof global.las);
global.las[0] = 0;
global.rt[0] = global.doBuild(1, 0, V);
part.rt[0] = part.doBuild(1, 0, V);
calc(1, 0, n);
for (int i = 1;i<= n;i++) {
if (!vis[i] && t[i] == 0) output[i]++;
printf("%d\n", output[i]);
}
return 0;
}
::::
到这里,略微卡常就可以过去了,但是这并不优美,你发现,可以将全局的贡献提出去,局部的负贡献不变,然后便可通过离线双指针加树状数组解决。
下面这个代码过的点的最大时间为
::::success[代码]
提交记录
#include <bits/stdc++.h>
using namespace std;
const int N = 5e5+10;
const int M = 1e6+10;
const int del = 5e5;
const int V = 1e6;
const int inf = 0x3f3f3f3f;
struct node {
int u, dep, n1, n2;
};
class BinaryIndexedTree {
private:
int c[V];
public:
int lowbit(int x) {
return x & -x;
}
void change(int x, int dx) {
while (x <= V) {
c[x] += dx;
x += lowbit(x);
}
}
int query(int x) {
int res = 0;
while (x) {
res += c[x];
x -= lowbit(x);
}
return res;
}
} pt;
int n, d;
int t[N];
int cnt, head[N];
int nxt[M], to[M];
int dep[N];
int now;
int siz[N], wight[N];
vector<node> vec;
bool vis[N];
int output[N];
void connect(int u, int v) {
cnt++, nxt[cnt] = head[u];
head[u] = cnt, to[cnt] = v;
}
void getCentroid(int u, int from, int &ctre) {
siz[u] = 1, wight[u] = 1;
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (v == from || vis[v])
continue;
getCentroid(v, u, ctre);
wight[u] = max(wight[u], siz[v]);
siz[u] += siz[v];
}
wight[u] = max(wight[u], now - siz[u]);
if (wight[u] <= now / 2) ctre = u;
}
void dfs1(int u, int from, int n1, int n2) {
dep[u] = dep[from] + 1; siz[u] = 1;
n1 = min(n1, t[u] - dep[u]);
n2 = min(n2, t[u] + dep[u]);
vec.push_back({u, dep[u], n1, n2});
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (v == from || vis[v])
continue;
dfs1(v, u, n1, n2);
siz[u] += siz[v];
}
}
void solve(int sign) {
if (vec.empty()) return ;
sort(vec.begin(), vec.end(), [](node x, node y) {
return x.dep < y.dep;
});
for (auto nd : vec) {
pt.change(nd.n1 + del, 1);
}
int j = vec.size()-1, res = 0;
for (int i = 0;i< vec.size();i++) {
while (j >= 0 && vec[i].dep + vec[j].dep > d) {
pt.change(vec[j].n1+del, -1), j--;
}
if (j < 0) break;
if (vec[i].n2 <= vec[i].dep) {
res = j + 1;
} else {
res = pt.query(vec[i].dep + del);
}
output[vec[i].u] += sign * res;
}
while (j >= 0) {
pt.change(vec[j].n1 + del, -1);
j--;
}
}
void calc(int u, int from, int num) {
int centre = u; now = num;
getCentroid(u, from, centre);
u = centre;
dep[0] = -1, dfs1(u, 0, inf, inf);
solve(1); vec.clear();
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (vis[v]) continue;
dfs1(v, u, t[u]-dep[u], t[u]+dep[u]);
solve(-1);
vec.clear();
}
vis[u] = 1;
for (int i = head[u];i;i=nxt[i]) {
int v = to[i];
if (vis[v]) continue;
calc(v, u, siz[v]);
}
}
signed main() {
cin.tie(0)->sync_with_stdio(false);
cin >> n >> d;
for (int i = 1;i<= n;i++) {
cin >> t[i];
} int u, v;
for (int i = 1;i< n;i++) {
cin >> u >> v;
connect(u, v);
connect(v, u);
}
calc(1, 0, n);
for (int i = 1;i<= n;i++) {
cout << output[i] << '\n';
}
return 0;
}
::::