题解:P15947 [JOI Final 2026] 集邮 5 / Collecting Stamps 5

· · 题解

一个非常点分治的好题。

题目就是说,让你求对于每个 u 来说,有多少个点 v(1 \le v \le v) 满足 uv 的路径上存在到达某个点的时间不小于题目给定的到达那个点的时间。

更公式化地说,你需要求有多少条路径 p_0,p_1, \dots ,p_m(m \le d),满足 \exists‌ i, t_{p_i} \le dist(p_0, p_i),然后将贡献计算到点 p_0 上。

考虑把不等式转化一下:

\begin{aligned} t_{p_i} &\le dist(p_0, p_i) \\ t_{p_i} &\le dep_{p_0} + dep_{p_i} - 2 \times dep_{lca(p_0, p_i)} \\ t_{p_i} - dep_{p_i} &\le dep_{p_0} - 2 \times dep_{lca(p_0, p_i)} \end{aligned}

既然转化成了含有最近公共祖先的式子,就不难考虑到点分治。

将式子拆开分成两个部分,一个部分是 p_0lca,另一个部分就是 lcap_m

两个部分的式子分别是 t_{p_i} + dep_{p_i} \le dep_{p_0}t_{p_i} - dep_{p_i} \le dep_{p_0} - 2 \times dep_{lca(p_0, p_i)}

这下就很明显了,对于一个分治中心 c,你考虑枚举 p_0,有两种情况,要么是在 p_0c 就已经满足了条件,这时候,你对于除了所在子树的点都是可以取到了(现在不考虑 d 的限制)。要么不满足,这时候,你就需要求满足第二种条件的数量了。

那如何判断呢?

简单地,其充要条件就是一条路径的最小值是否满足条件就行了。

证明是简单的。充分的证明是显然的,你一条路径中值最小的点满足 \le 的条件,那这条路径就一定就是满足条件了的呀。必要的证明就是,你最小的值都不满足 \le 的条件,难道还有其他的值满足?

然后你通过总体满足条件的数量减去所在子树满足条件的数量的写法,发现只得到了 41 分,这是因为没有 d 的限制。

加上了 d 的限制后,你发现就是一个二维偏序,用两棵主席树便可以完成。

然后就做完了。

提前说一下,这个过的最大时间为 4.36 秒。

::::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;
} 

::::

到这里,略微卡常就可以过去了,但是这并不优美,你发现,可以将全局的贡献提出去,局部的负贡献不变,然后便可通过离线双指针加树状数组解决。

下面这个代码过的点的最大时间为 1.68 秒。

::::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;
} 

::::