1 条题解
-
0
#include <bits/stdc++.h> using namespace std; #define lc (p << 1) #define rc (p << 1 | 1) #define mid (tr[p].l + tr[p].r) / 2 const int N = 1e6 + 10; vector<pair<int, int>> G[N]; int fa[N], son[N], dep[N], f[N][20], D, siz[N], a[N]; void dfs1(int x, int ff) { fa[x] = ff; dep[x] = dep[ff] + 1; siz[x] = 1; f[x][0] = ff; for (int i = 1; i <= D; i++) f[x][i] = f[f[x][i - 1]][i - 1]; for (auto i : G[x]) { int y = i.first, c = i.second; if (y == ff) continue; dfs1(y, x); a[y] = c; siz[x] += siz[y]; if (siz[son[x]] < siz[y]) son[x] = y; } } int tsp, dfn[N], top[N], ys[N]; void dfs2(int x, int tp) { dfn[x] = ++tsp; ys[tsp] = x; top[x] = tp; if (son[x] != 0) dfs2(son[x], tp); for (auto i : G[x]) { int y = i.first, c = i.second; if (y != fa[x] && y != son[x]) dfs2(y, y); } } struct trnode { int l, r, c; } tr[N << 2]; void upd(int p) { tr[p].c = tr[lc].c + tr[rc].c; } void bt(int p, int l, int r) { tr[p] = {l, r, 0}; if (l == r) tr[p].c = a[ys[l]]; else { bt(lc, l, mid); bt(rc, mid + 1, r); upd(p); } } int query(int p, int l, int r) { if (l <= tr[p].l && tr[p].r <= r) return tr[p].c; int res = 0; if (l <= mid) res += query(lc, l, r); if (mid < r) res += query(rc, l, r); return res; } int solve(int x, int y) { int res = 0; while (top[x] != top[y]) { if (dep[top[x]] < dep[top[y]]) swap(x, y); res += query(1, dfn[top[x]], dfn[x]); x = fa[top[x]]; } if (dep[x] > dep[y]) swap(x, y); if (x != y) res += query(1, dfn[x] + 1, dfn[y]); return res; } int lca(int x, int y) { if (dep[x] < dep[y]) swap(x, y); for (int i = D; i >= 0; i--) if (dep[f[x][i]] >= dep[y]) x = f[x][i]; if (x == y) return x; for (int i = D; i >= 0; i--) if (f[x][i] != f[y][i]) x = f[x][i], y = f[y][i]; return f[x][0]; } int path(int x, int y, int k) { int p = lca(x, y); if (k <= dep[x] - dep[p] + 1) { k--; for (int i = D; i >= 0; i--) if (k >= (1 << i)) x = f[x][i], k -= (1 << i); return x; } else { k = (dep[x] + dep[y] - 2 * dep[p] + 1) - k + 1; k--; for (int i = D; i >= 0; i--) if (k >= (1 << i)) y = f[y][i], k -= (1 << i); return y; } } int main() { int T; scanf("%d", &T); while (T--) { int n; scanf("%d", &n); memset(G, 0, sizeof(G)); for (int i = 1, x, y, c; i < n; i++) { scanf("%d%d%d", &x, &y, &c); G[x].push_back({y, c}); G[y].push_back({x, c}); } fa[0] = dep[0] = siz[0] = 0; memset(son, 0, sizeof(son)); a[1] = 0; D = log2(n); dfs1(1, 0); tsp = 0; dfs2(1, 1); bt(1, 1, tsp); char s[10]; while (scanf("%s", s) != EOF && s[1] != 'O') { if (s[0] == 'D') { int x, y; scanf("%d%d", &x, &y); printf("%d\n", solve(x, y)); } else { int x, y, k; scanf("%d%d%d", &x, &y, &k); printf("%d\n", path(x, y, k)); } } } return 0; }
- 1
信息
- ID
- 546
- 时间
- 1000ms
- 内存
- 2048MiB
- 难度
- 8
- 标签
- 递交数
- 76
- 已通过
- 11
- 上传者