题解:P9983 [USACO23DEC] Cowntact Tracing P

· · 题解

该题解对做法正确性提供大量解释,写得比较详细。本题并没有那么难。

【sol】

考虑选点的策略。最终每个黑点都要被覆盖一次,我们不妨以深度从深到浅考虑每个点。

对于黑色节点 u,在之前的考虑中他未被覆盖,必须要操作一次 u,满足 \operatorname{dis}(u,v)\leq d,且 v 到最近的白色点距离 d_1>d,这样操作 v 才不会让白色点被覆盖到。

我们可以多源 bfs 去预处理 d_1,所以对于每次询问可以快速找到可以操作的节点,设这些节点构成集合 cand。问题转化为,对于所有 \operatorname{dis }(u,v)\leq d,v\in cand,找一个最优的操作。

定义节点 vd -邻域为 D_v,则 v_1v_2 更优定义为:D_{v_2}\cap S\subseteq D_{v_1}\cap S,其中 S 为所有未被覆盖的黑色点构成的集合。

对于 \operatorname{dep}(v_1)<\operatorname{dep}(v_2),v_1,v_2\in cand,试比较 v_1,v_2 哪个更优。设 l_1=\operatorname{lca}(u,v_1),l_2=\operatorname{lca}(u,v_2)

因为 u 是未被覆盖过的最深的点,v_1,v_2 能覆盖 u,就能覆盖 l_1,l_2 子树内所有未被覆盖的点。在这方面他们是同等优秀的。对于 l_1l_2 的祖先,我们不关心 l_1,l_2 的位置关系,因为 \operatorname{dep}(v_1)<\operatorname{dep}(v_2),所以 v_1 在祖先上的覆盖是优于 v_2 的,故 v_1 严格优于 v_2

现在就转化为了如何对 u 找到 \operatorname{dep} 最小的 v\in cand,\operatorname{dis}(u,v)\leq d

拆式子。先 bfs 找出每个节点 x 最近的 v\in cand,设 mn_x 为其距离。

所以将 v 贡献折算到其与 u 广义的 \operatorname{lca} 上,设其为 l,则式子重写为 \operatorname{dis}(u,v)=\operatorname{dep}_u-\operatorname{dep}_l+mn_l\leq d,即 \operatorname{dep}_l+d-mn_l\leq\operatorname{dep}_u。可以注意到式子左侧是 l 的固有属性,可用倍增找到最浅的 l,设 val_x=\operatorname{dep}_x+d-mn_x

但是迎来了一个问题,最浅的 l 满足条件其对应的 v 也最浅吗?这需要我们说明一下。

假设最浅的符合条件的 v\in candv_1,其在 l_1 处符合条件,并且对于 l_1 每向上移动一步 val 对应减二,所以整体上讲,找到的最浅的 ll_1 的一个祖先。

若最浅的 l 对应的 v 不在 l 子树内,则注意到 l 向上移动一步后 \operatorname{dis} 减小,\operatorname{dep} 也减小,所以 val_{fa_l}=val_l,同样符合条件,与 l 是最浅的冲突。所以 l 对应的 v 一定在其子树内。而如果 l=1l 不存在父亲,那么显然 v 同样在 1 的子树内。

\operatorname{dis}(v,l) 可改写为 \operatorname{dep}_v-\operatorname{dep}_l。又因为 v_1 是最浅的,则显然此时 v_1 最优。所以我们论证了最浅的 l 满足 val_l\leq dep_u,其对应的 v 一定是全局最浅的 v

那么这题到这里就结束了,倍增找到 l。同时我们还要对 vd -邻域做覆盖操作,并询问一个节点的状态,这个可以用点分树去做。

最后复杂度是 O(nq\log n),需要实现一个欧拉序 \operatorname{lca}O(1) 求两点距离。

【code】

#include<bits/stdc++.h>
using namespace std;

const int nr = 2e5 + 10;
const int lr = 21;
const int inf = 2e9;
int lg[nr << 1], n, q, d, a[nr], dep[nr], anc[nr][lr], id[nr], dis1[nr], dis2[nr], pos[nr], val[nr];
int sz[nr], rt, rtmx, SZ, fa[nr], len[nr]; bool vis[nr], secc;
vector<int> adj[nr];

namespace LCA
{
    int st[nr << 1][lr], euler[nr << 1]; 
    int etot;
    void dfs(int x, int ft)
    {
        etot++, euler[x] = etot, st[etot][0] = x;
        dep[x] = dep[ft] + 1, anc[x][0] = ft;
        for (int i = 1; i <= lg[dep[x]]; i++)
            anc[x][i] = anc[anc[x][i - 1]][i - 1];
        for (int i = 0; i < adj[x].size(); i++)
            if (adj[x][i] != ft) dfs(adj[x][i], x), st[++etot][0] = x;
    }
    void init()
    {
        etot = 0, dfs(1, 0);
        for (int len = 1; 1 << len <= etot; len++)
            for (int i = 1; i + (1 << len) - 1 <= etot; i++)
                st[i][len] = dep[st[i][len - 1]] < dep[st[i + (1 << len - 1)][len - 1]] ? st[i][len - 1] : st[i + (1 << len - 1)][len - 1]; 
    }
    int query(int l, int r)
    {
        if (l > r) swap(l, r); 
        int k = lg[r - l + 1];
        return dep[st[l][k]] < dep[st[r - (1 << k) + 1][k]] ? st[l][k] : st[r - (1 << k) + 1][k];
    }
    int lca(int x, int y)
    { 
        int ex = euler[x], ey = euler[y]; 
        return query(ex, ey);
    }
    int dis(int x, int y)
    {
        return dep[x] + dep[y] - 2 * dep[lca(x, y)];
    }
}

void initrt(int x, int ft)
{
    int mx = 0; sz[x] = 1;
    for (int i = 0; i < adj[x].size(); i++)
    {
        int v = adj[x][i];
        if (v == ft || vis[v]) continue;
        initrt(v, x);
        sz[x] += sz[v], mx = max(mx, sz[v]);
    }
    mx = max(mx, SZ - sz[x]);
    if (mx < rtmx) rt = x, rtmx = mx;
}

int initsz(int x, int ft)
{
    int res = 1;
    for (int i = 0; i < adj[x].size(); i++)
    {
        int v = adj[x][i];
        if (v == ft || vis[v]) continue;
        res += initsz(v, x);
    }
    return res;
}

void init(int x)
{
    vis[x] = true;
    for (int i = 0; i < adj[x].size(); i++)
    {
        int v = adj[x][i];
        if (vis[v]) continue;
        SZ = initsz(v, x), rt = 0, rtmx = inf;
        initrt(v, x), fa[rt] = x, init(rt);
    }
}

void bfs1()
{
    for (int i = 1; i <= n; i++) dis1[i] = inf;
    queue<int> q;
    for (int i = 1; i <= n; i++) if (!a[i]) q.push(i), dis1[i] = 0;
    while (!q.empty())
    {
        int u = q.front(); q.pop();
        for (int i = 0; i < adj[u].size(); i++)
            if (dis1[adj[u][i]] == inf) dis1[adj[u][i]] = dis1[u] + 1, q.push(adj[u][i]);
    }
}

void bfs2()
{
    for (int i = 1; i <= n; i++) dis2[i] = inf;
    queue<int> q;
    for (int i = 1; i <= n; i++) if (dis1[i] > d) q.push(i), dis2[i] = 0, pos[i] = i;
    if (q.empty()) { secc = false; return; }
    while (!q.empty())
    {
        int u = q.front(); q.pop();
        for (int i = 0; i < adj[u].size(); i++)
            if (dis2[adj[u][i]] == inf) dis2[adj[u][i]] = dis2[u] + 1, pos[adj[u][i]] = pos[u], q.push(adj[u][i]);
    }
}

bool check(int x)
{
    int cur = x;
    while (cur)
    {
        if (len[cur] >= LCA::dis(x, cur)) return true;
        cur = fa[cur];
    }
    return false;
}

void ban(int x, int y)
{
    int cur = x;
    while (cur)
    {
        int now = d - LCA::dis(x, cur);
        if (now >= 0) len[cur] = max(len[cur], now);
        cur = fa[cur];
    }
}

int find(int x)
{
    for (int i = lg[dep[x]], lim = dep[x]; i >= 0; i--)
        if (val[anc[x][i]] >= lim) x = anc[x][i];
    return x;
}

int main()
{
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0); 
    lg[0] = -1;
    for (int i = 1; i < nr << 1; i++) lg[i] = lg[i >> 1] + 1;
    cin >> n;
    for (int i = 1; i <= n; i++) { char c; cin >> c; a[i] = c - '0'; }
    for (int i = 1; i < n; i++)
    {
        int u, v; cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    SZ = n, rt = 0, rtmx = inf, init(1);
    bfs1(), LCA::init();
    for (int i = 1; i <= n; i++) id[i] = i;
    sort(id + 1, id + n + 1, [&](int x, int y) { return dep[x] > dep[y]; });
    cin >> q;
    while (q--)
    {
        cin >> d;
        secc = true; bfs2();
        if (!secc) { cout << -1 << '\n'; continue; }
        for (int i = 1; i <= n; i++) val[i] = dep[i] + d - dis2[i], len[i] = -1;
        int res = 0;
        for (int i = 1; i <= n; i++)
        {
            int x = id[i];
            if (!a[x] || check(x)) continue;
            if (val[x] < dep[x]) { res = -1; break; }
            ban(pos[find(x)], x), res++; 
        }
        cout << res << '\n'; 
    }
    return 0;
}