【题解】P15952 [ICPC 2018 Jakarta R] Rotating Gears

P15952 [ICPC 2018 Jakarta R] Rotating Gears 题解

题意

给定一棵 \(N\) 个节点的树,每个节点是一个齿轮,初始所有齿轮箭头指向 \(0\) 度。每个齿轮有两种状态:在板上或被取出。有三种操作:1 x 取出齿轮 \(x\)2 x 放回齿轮 \(x\)(保持取出时的角度);3 x α 将齿轮 \(x\) 顺时针旋转 \(\alpha\) 度。由于齿轮相互接触,当齿轮 \(x\) 被旋转时,\(x\) 所在的连通块(仅由当前在板上的齿轮组成)内所有齿轮都会随之旋转,其中与 \(x\) 深度奇偶相同的齿轮顺时针转 \(\alpha\),与 \(x\) 深度奇偶不同的齿轮逆时针转 \(\alpha\)。每次操作 3 需要输出能量消耗,定义为旋转的齿轮数量乘以 \(\alpha\);最后输出所有齿轮最终箭头的顺时针角度之和,每个角度先对 \(360\) 取模再求和。

分析

直接模拟每次从 \(x\) 出发 BFS 找连通块并逐个修改角度,单次最坏 \(O(N)\),总复杂度 \(O(NQ)\),无法通过 \(N, Q \le 10^5\)。我们需要一个能快速定位连通块并批量修改角度的做法。

很简单的一个想法是:定义 \(tag_u\) 表示一个节点是否被取下,\(0\) 则没有被取下,\(1\) 则被取下。

关键观察是:不能只看 \(tag_u\)(是否被取出)来判断两个节点是否联通。例如链 \(1-2-3-4\) 取出节点 \(2\),节点 \(1\)\(3\)\(tag\) 都是 \(0\),但它们并不联通,因为中间隔着被取出的 \(2\)。正确的量是看根到节点的 \(tag\) 前缀和是否相等:

\[S_u = \sum_{v \in \text{path}(1 \to u)} tag_v \]

即根到 \(u\) 路径上被取出齿轮的数量。初始所有 \(tag = 0\),所以 \(S_u = 0\)。由于 \(tag\) 非负,\(S\) 沿子树单调不降。若 \(u\)\(v\) 的祖先且两者都在板上,则 \(u\)\(v\) 联通当且仅当 \(S_u = S_v\),因为路径 \(u \to v\)\(tag\) 之和为 \(S_v - S_u\),若为 \(0\) 则路径上没有被取出节点,反之若联通则路径上没有取出节点、和为 \(0\)

我们设 \(x\) 是在板上的节点,它所在连通块的顶 \(t\)\(x\) 向上能找到的最浅的、在板上且 \(S_t = S_x\) 的祖先。而连通块大小为

\[sz = |\{v \in subtree(t) \mid S_v = S_t\}| \]

由于 \(S\) 在子树内单调不降,\(S_t\) 恰是 \(subtree(t)\)\(S\) 的最小值,所以

\[sz = |\{v \in subtree(t) \mid S_v = \min_{w \in subtree(t)} S_w\}| \]

这提示我们用线段树维护区间最小值和最小值出现次数。

接下来看操作对 \(S\) 的影响。操作 1 x 使 \(tag_x\)\(0\)\(1\),对 \(x\) 子树内所有节点,根路径都经过 \(x\),所以 \(S\) 都加 \(1\);操作 2 x 使 \(tag_x\)\(1\)\(0\),对 \(x\) 子树内所有节点 \(S\)\(1\)。用 DFS 序把子树变成区间 \([in_x, out_x]\),于是操作 1/2 就是区间加/减 \(1\)

对于角度,一次 3 x α 对连通块内节点 \(v\) 的影响是:若 \(dep_v \operatorname{mod} 2 = dep_x \operatorname{mod} 2\),角度加 \(\alpha\);否则角度减 \(\alpha\)。我们引入带符号角度,规定顺时针为正、逆时针为负,则影响统一为加法,加数为 \(\alpha\)\(-\alpha\)。最终计算时取 \(\delta_u = ((ang_u \bmod 360) + 360) \bmod 360\),负值会自动转成对应的顺时针角度。

基于 DFS 序建一棵线段树 A,每个节点维护:\(mn\) 表示区间内 \(S\) 的最小值,\(c\) 表示区间内 \(S = mn\) 的节点个数,\(lz\) 表示区间 \(S\) 的加法懒标记,\(la[2]\) 表示对区间内 \(S = mn\) 的节点、深度奇偶为 \(0/1\) 的带符号角度增量。区间加 add 是标准线段树区间加;qry 返回 (mn, c)upd 对完全覆盖的节点,若 mn == val 则直接对该节点的 la 打标记:la[p] += αla[1-p] -= α,若 mn > val 则说明该区间没有目标节点、直接返回。push 的关键逻辑是:父节点向下推 lzla 时,只有子节点的 mn == 父节点旧 mn 才把角度懒标记传下去,因为只有这些子节点是目标节点。

为了快速定位连通块顶,用树链剖分把 \(x \to root\) 路径拆成若干重链区间,用 set<int> 维护所有 tag = 1 节点的 DFS 序。查找 \(x\) 路径上最深的 \(tag = 1\) 节点 \(p\) 时,沿重链向上,在当前链区间 \([in_{tp}, in_x]\) 中找最大的 \(tg = 1\) 位置,若找到则对应节点就是 \(p\),否则跳到 \(f_{tp}\) 继续。找到 \(p\) 后,若 \(p = 0\) 说明 \(x\) 与根连通、\(t = 1\);否则用倍增从 \(x\) 向上跳到深度刚好比 \(p\)\(1\) 的位置,即 \(p\) 在路径上的那个儿子,即为 \(t\)

算法流程

\(1\) 为根 DFS 预处理 f, dep, sz, sn, tp, in, out, rv 和倍增数组 up[k][u]。建线段树 A,初始化 mn = 0c = 区间长度、懒标记全为 \(0\)。处理每次操作:1 xtg[x] = 1pos.insert(in[x])seg.add(in[x], out[x], +1)2 xtg[x] = 0pos.erase(in[x])seg.add(in[x], out[x], -1)3 x α 时找 \(p\) 为路径 \(x \to root\) 上最深的 tg = 1 节点,\(t = 1\)(若 \(p = 0\))或 gcp(p, x),查询 qry(in[t], out[t]) 得到 mnsz,输出 sz * α,并调用 upd(in[t], out[t], mn, dep[x] % 2, α)。最后递归下推所有懒标记到叶子,得到每个节点的带符号角度 angle[u],对每个节点取模 \(360\)(保证非负)后累加求和输出。

复杂度

操作 1/2 为 \(O(\log N)\),操作 3 找顶为 \(O(\log^2 N)\)、查询与角度更新各为 \(O(\log N)\),最终下推为 \(O(N)\),总时间复杂度 \(O(Q \log^2 N)\),空间 \(O(N)\),在 \(N, Q \le 10^5\) 下运行时间约 \(1\) 秒以内。

代码

#include <bits/stdc++.h>
#define lc u << 1
#define rc u << 1 | 1
using namespace std;
const int N = 1e5 + 5;
const int LOG = 18;
int n, q;
vector<int> g[N];
int f[N], dep[N], sz[N], sn[N], tp[N];
int in[N], out[N], rv[N], tm_;
int up[LOG][N];
int tg[N];
set<int> pos;
struct SGT{
    int mn[N << 2], c[N << 2];
    int lz[N << 2], la[N << 2][2];
    inline void build(int u, int l, int r){
        mn[u] = 0;
        c[u] = r - l + 1;
        lz[u] = 0;
        la[u][0] = la[u][1] = 0;
        if (l == r) return;
        int mid = (l + r) >> 1;
        build(lc, l, mid), build(rc, mid + 1, r);
        return;
    }
    inline void pull(int u){
        mn[u] = min(mn[lc], mn[rc]);
        c[u] = 0;
        if (mn[lc] == mn[u]) c[u] += c[lc];
        if (mn[rc] == mn[u]) c[u] += c[rc];
    }
    inline void push(int u){
        if (lz[u] == 0 && la[u][0] == 0 && la[u][1] == 0) return;
        int old = mn[u] - lz[u];
        if (mn[lc] == old) {
            la[lc][0] += la[u][0];
            la[lc][1] += la[u][1];
        }
        mn[lc] += lz[u];
        lz[lc] += lz[u];
        if (mn[rc] == old) {
            la[rc][0] += la[u][0];
            la[rc][1] += la[u][1];
        }
        mn[rc] += lz[u];
        lz[rc] += lz[u];
        lz[u] = 0;
        la[u][0] = la[u][1] = 0;
    }
    inline void add(int u, int l, int r, int ql, int qr, int d){
        if (ql <= l && r <= qr){
            mn[u] += d, lz[u] += d;
            return;
        }
        push(u);
        int mid = (l + r) >> 1;
        if(ql <= mid)
            add(lc, l, mid, ql, qr, d);
        if(qr > mid)
            add(rc, mid + 1, r, ql, qr, d);
        pull(u);
        return;
    }
    inline pair<int,int> qry(int u, int l, int r, int ql, int qr) {
        if(ql <= l && r <= qr)
            return {mn[u], c[u]};
        push(u);
        int mid = (l + r) >> 1;
        if(qr <= mid)
            return qry(lc, l, mid, ql, qr);
        if(ql > mid)
            return qry(rc, mid + 1, r, ql, qr);
        auto L = qry(lc, l, mid, ql, qr);
        auto R = qry(rc, mid + 1, r, ql, qr);
        int m = min(L.first, R.first), cc = 0;
        if(L.first == m)
            cc += L.second;
        if(R.first == m)
            cc += R.second;
        return {m, cc};
    }
    inline void upd(int u, int l, int r, int ql, int qr, int val, int p, int a){
        if(ql <= l && r <= qr){
            if(mn[u] > val)
                return;
            if(mn[u] == val){
                la[u][p] += a, la[u][1 - p] -= a;
                return;
            }
        }
        push(u);
        int mid = (l + r) >> 1;
        if(ql <= mid)
            upd(lc, l, mid, ql, qr, val, p, a);
        if(qr > mid)
            upd(rc, mid + 1, r, ql, qr, val, p, a);
        pull(u);
        return;
    }
    inline void pd(int u, int l, int r, vector<int>& ang){
        if(l == r){
            int nd = rv[l];
            ang[nd] = la[u][dep[nd] % 2];
            return;
        }
        push(u);
        int mid = (l + r) >> 1;
        pd(lc, l, mid, ang), pd(rc, mid + 1, r, ang);
        return;
    }
}seg;
inline void d1(int u, int p){
    f[u] = p;
    dep[u] = dep[p] + 1;
    sz[u] = 1;
    sn[u] = 0;
    for(int v : g[u]){
        if(v == p)
            continue;
        d1(v, u);
        sz[u] += sz[v];
        if(sn[u] == 0 || sz[v] > sz[sn[u]])
            sn[u] = v;
    }
    return;
}
inline void d2(int u, int t){
    tp[u] = t;
    in[u] = ++tm_;
    rv[tm_] = u;
    if(sn[u])
        d2(sn[u], t);
    for(int v : g[u]){
        if(v == f[u] || v == sn[u])
            continue;
        d2(v, v);
    }
    out[u] = tm_;
    return;
}
inline int fnd(int x){
    while(x != 0){
        int t = tp[x];
        auto it = pos.upper_bound(in[x]);
        if(it != pos.begin()){
            --it;
            if(*it >= in[t])
                return rv[*it];
        }
        x = f[t];
    }
    return 0;
}
inline int gcp(int p, int x){
    for(int k = LOG - 1; k >= 0; k--){
        if(up[k][x] != 0 && dep[up[k][x]] > dep[p])
            x = up[k][x];
    }
    return x;
}
inline void read(int& x){
    int s = 0, w = 1;
    char ch = getchar();
    while(!isdigit(ch)){
        w = ch == '-' ? -1 : 1;
        ch = getchar();
    }
    while(isdigit(ch)){
        s = s * 10 + ch - '0';
        ch = getchar();
    }
    x = s * w;
    return;
}
inline void write(long long x){
    if(x < 0) x = -x, putchar('-');
    if(x > 9) write(x / 10);
    putchar(x % 10 + '0');
    return;
}
int main(){
    read(n);
    for (int i = 0; i < n - 1; i++) {
        int u, v;
        read(u), read(v);
        g[u].push_back(v), g[v].push_back(u);
    }
    dep[0] = -1;
    d1(1, 0);
    tm_ = 0;
    d2(1, 1);
    for(int i = 1; i <= n; i++)
        up[0][i] = f[i];
    for(int k = 1; k < LOG; k++)
        for(int i = 1; i <= n; i++)
            up[k][i] = up[k - 1][up[k - 1][i]];
    seg.build(1, 1, n);
    memset(tg, 0, sizeof(tg));
    read(q);
    while (q--) {
        int op;
        read(op);
        if(op == 1){
            int x;
            read(x);
            if(!tg[x]){
                tg[x] = 1;
                pos.insert(in[x]);
                seg.add(1, 1, n, in[x], out[x], 1);
            }
        }else if(op == 2){
            int x;
            read(x);
            if(tg[x]){
                tg[x] = 0;
                pos.erase(in[x]);
                seg.add(1, 1, n, in[x], out[x], -1);
            }
        }else{
            int x, a;
            read(x), read(a);
            int p = fnd(x);
            int t = (p == 0) ? 1 : gcp(p, x);
            auto res = seg.qry(1, 1, n, in[t], out[t]);
            int m = res.first, s = res.second;
            write(1LL * s * a), putchar('\n');
            seg.upd(1, 1, n, in[t], out[t], m, dep[x] % 2, a);
        }
    }
    vector<int> ang(N + 1, 0);
    seg.pd(1, 1, n, ang);
    long long tot = 0;
    for(int i = 1; i <= n; i++){
        int a = ang[i] % 360;
        if(a < 0)
            a += 360;
        tot += a;
    }
    write(tot), putchar('\n');
    return 0;
}
posted @ 2026-09-10 20:05  Hyazinth  阅读(4)  评论(0)    收藏  举报