【题解】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\) 前缀和是否相等:
即根到 \(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\) 的祖先。而连通块大小为
由于 \(S\) 在子树内单调不降,\(S_t\) 恰是 \(subtree(t)\) 中 \(S\) 的最小值,所以
这提示我们用线段树维护区间最小值和最小值出现次数。
接下来看操作对 \(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 的关键逻辑是:父节点向下推 lz 和 la 时,只有子节点的 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 = 0、c = 区间长度、懒标记全为 \(0\)。处理每次操作:1 x 时 tg[x] = 1、pos.insert(in[x])、seg.add(in[x], out[x], +1);2 x 时 tg[x] = 0、pos.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]) 得到 mn 和 sz,输出 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;
}

浙公网安备 33010602011771号