#P3443. 第3题-集市之路
-
1000ms
Tried: 33
Accepted: 3
Difficulty: 10
所属公司 :
美团
时间 :2025年8月23日-开发岗
算法与标签>树
第3题-集市之路
思路总览\n\n核心目标:在动态切换集市状态的同时,高效回答“到所有开设集市的城市的距离之和”。直接维护会超时,因此使用点分治配合分层距离与分组前缀量维护。\n\n### 关键记号\n\n* 用点分治把树分成若干层重心。对每个原树节点 u,记录它到每一层重心的距离序列:\n\n
\n\n 其中 ci 是某一层的重心(从“最深”那层到“最顶”那层),di=dist(u,ci)。\n* 对每个重心 c 维护两类量:\n\n 1. 全体开设集市的城市在 c 处的聚合量\n\n
\n\n 2. “沿着某个子分治方向”的子聚合量(用“下层重心”作键)\n\n
\n\n 这里的 p 是 c 的下层重心(即点分治树上 c 的某个孩子),它唯一代表“从 c 走向某个连通分量的方向”。\n\n### 更新(切换集市点 u,令 Delta=+1 表示开设集市、Delta=−1 表示关闭集市)\n\n沿 u 的重心链自下而上遍历:\n设当前重心为 c,且与之下一层的重心为 p(若不存在则忽略第二项),并记 d=dist(u,c)。\n做\n\n
\n\n若存在p,再做\n\n
\n\n### 查询(求 Sv)\n\n同样沿 v 的重心链自下而上遍历:\n设当前重心为 c,下层重心为 p(若不存在忽略扣除项),d=dist(v,c)。\n总和累加:\n\n
\n\n为避免把“与 v 同处 c 的同一子方向”的集市城市重复计入,需要扣除:\n\n
\n\n处理完整条重心链后,所得即为 Sv。\n\n\n# C++ \n\ncpp\n#include <bits/stdc++.h>\nusing namespace std;\n\n// --------------- 全局结构 ---------------\nstruct Edge { int to; int w; };\nconst int MAXN = 200000 + 5;\n\nint n, q;\nvector<Edge> g[MAXN];\n\n// 点分治需要的标记与数据\nbool removed_[MAXN];\nint sz[MAXN];\nint cpar[MAXN]; // 点分治树上的父重心(根为 -1)\n\n// 对每个原树节点:其重心链(从上到下/从顶到底存储),以及到这些重心的距离\nvector<int> cpath[MAXN];\nvector<long long> cdist[MAXN];\n\n// 统计量\nlong long totCnt[MAXN], totDistSum[MAXN];\n// 每个重心 c 的子方向统计量:使用 unordered_map<下层重心, 值>\nunordered_map<int, long long> subCnt[MAXN], subDistSum[MAXN];\n\n// 初始集市状态\nvector<int> initRed;\nvector<char> isRed;\n\n// --------------- 点分治(递归版) ---------------\nint calcSize(int u, int p) {\n sz[u] = 1;\n for (auto e : g[u]) {\n int v = e.to;\n if (v == p || removed_[v]) continue;\n sz[u] += calcSize(v, u);\n }\n return sz[u];\n}\n\nint findCentroid(int u, int p, int tot) {\n for (auto e : g[u]) {\n int v = e.to;\n if (v == p || removed_[v]) continue;\n if (sz[v] > tot / 2) return findCentroid(v, u, tot);\n }\n return u;\n}\n\nvoid collect(int u, int p, int cen, long long dist) {\n // 把 cen 挂到 u 的重心链末尾,并记录距离\n cpath[u].push_back(cen);\n cdist[u].push_back(dist);\n for (auto e : g[u]) {\n int v = e.to;\n if (v == p || removed_[v]) continue;\n collect(v, u, cen, dist + e.w);\n }\n}\n\nvoid decompose(int entry, int parent) {\n int tot = calcSize(entry, -1);\n int cen = findCentroid(entry, -1, tot);\n // 先收集 cen 覆盖的本组件所有点到 cen 的距离\n collect(cen, -1, cen, 0);\n cpar[cen] = parent;\n removed_[cen] = true;\n // 递归处理每个未移除的相邻分量\n for (auto e : g[cen]) {\n int v = e.to;\n if (!removed_[v]) {\n decompose(v, cen);\n }\n }\n // 不必恢复 removed_[cen]\n}\n\n// --------------- 维护(切换 / 查询) ---------------\nvoid apply_update(int u, int delta) {\n // delta = +1 表示设为开设集市;-1 表示关闭集市\n int prev = -1;\n int m = (int)cpath[u].size();\n for (int i = m - 1; i >= 0; --i) {\n int c = cpath[u][i];\n long long d = cdist[u][i];\n totCnt[c] += delta;\n totDistSum[c] += 1LL * delta * d;\n if (prev != -1) {\n subCnt[c][prev] += delta;\n subDistSum[c][prev] += 1LL * delta * d;\n }\n prev = c;\n }\n}\n\nlong long query_sum(int u) {\n long long ans = 0;\n int prev = -1;\n int m = (int)cpath[u].size();\n for (int i = m - 1; i >= 0; --i) {\n int c = cpath[u][i];\n long long d = cdist[u][i];\n ans += totDistSum[c] + totCnt[c] * d;\n if (prev != -1) {\n auto it1 = subDistSum[c].find(prev);\n auto it2 = subCnt[c].find(prev);\n long long sd = (it1 == subDistSum[c].end() ? 0LL : it1->second);\n long long sc = (it2 == subCnt[c].end() ? 0LL : it2->second);\n ans -= sd + sc * d;\n }\n prev = c;\n }\n return ans;\n}\n\n// --------------- 主过程 ---------------\nint main() {\n ios::sync_with_stdio(false);\n cin.tie(nullptr);\n\n cin >> n >> q;\n initRed.resize(n + 1);\n isRed.assign(n + 1, 0);\n for (int i = 1; i <= n; ++i) {\n cin >> initRed[i];\n isRed[i] = (initRed[i] ? 1 : 0);\n }\n for (int i = 0; i < n - 1; ++i) {\n int u, v, w;\n cin >> u >> v >> w;\n g[u].push_back({v, w});\n g[v].push_back({u, w});\n }\n\n // 点分治建树 + 预处理各点到各层重心距离\n decompose(1, -1);\n\n // 根据初始集市状态批量加入\n for (int i = 1; i <= n; ++i) {\n if (isRed[i]) apply_update(i, +1);\n }\n\n // 在线处理操作\n for (int i = 0; i < q; ++i) {\n int t, v;\n cin >> t >> v;\n if (t == 1) {\n if (isRed[v]) {\n apply_update(v, -1);\n isRed[v] = 0;\n } else {\n apply_update(v, +1);\n isRed[v] = 1;\n }\n } else {\n cout << query_sum(v) << "\n";\n }\n }\n return 0;\n}\n\n# Python \n\npython\nimport sys\nsys.setrecursionlimit(1 << 25)\ndata = sys.stdin.buffer.read().split()\nit = iter(data)\ndef ni(): return int(next(it))\n\nn, q = ni(), ni()\ninit = [0]*(n+1)\nis_red = [0]*(n+1)\nfor i in range(1, n+1):\n init[i] = ni()\n is_red[i] = 1 if init[i] == 1 else 0\n\ng = [[] for _ in range(n+1)]\nfor _ in range(n-1):\n u, v, w = ni(), ni(), ni()\n g[u].append((v, w))\n g[v].append((u, w))\n\nremoved = [False]*(n+1)\nsz = [0]*(n+1)\ncpar = [-1]*(n+1)\n\n# 每个点到各层重心的链及距离(从上到下存,使用时倒序遍历)\ncpath = [[] for _ in range(n+1)]\ncdist = [[] for _ in range(n+1)]\n\ntotCnt = [0]*(n+1)\ntotDist = [0]*(n+1)\nfrom collections import defaultdict\nsubCnt = [defaultdict(int) for _ in range(n+1)]\nsubDist = [defaultdict(int) for _ in range(n+1)]\n\ndef calc_size(u, p):\n sz[u] = 1\n for v, w in g[u]:\n if v == p or removed[v]:\n continue\n calc_size(v, u)\n sz[u] += sz[v]\n\ndef find_centroid(u, p, tot):\n for v, w in g[u]:\n if v == p or removed[v]:\n continue\n if sz[v] > tot // 2:\n return find_centroid(v, u, tot)\n return u\n\ndef collect(u, p, cen, dist):\n # 把当前重心 cen 附加到 u 的重心链\n cpath[u].append(cen)\n cdist[u].append(dist)\n for v, w in g[u]:\n if v == p or removed[v]:\n continue\n collect(v, u, cen, dist + w)\n\ndef decompose(entry, parent):\n calc_size(entry, -1)\n cen = find_centroid(entry, -1, sz[entry])\n # 先收集距离\n collect(cen, -1, cen, 0)\n cpar[cen] = parent\n removed[cen] = True\n for v, w in g[cen]:\n if not removed[v]:\n decompose(v, cen)\n\ndef apply_update(u, delta):\n prev = -1\n m = len(cpath[u])\n for i in range(m-1, -1, -1):\n c = cpath[u][i]\n d = cdist[u][i]\n totCnt[c] += delta\n totDist[c] += delta * d\n if prev != -1:\n subCnt[c][prev] += delta\n subDist[c][prev] += delta * d\n prev = c\n\ndef query_sum(u):\n ans = 0\n prev = -1\n m = len(cpath[u])\n for i in range(m-1, -1, -1):\n c = cpath[u][i]\n d = cdist[u][i]\n ans += totDist[c] + totCnt[c] * d\n if prev != -1:\n ans -= subDist[c].get(prev, 0) + subCnt[c].get(prev, 0) * d\n prev = c\n return ans\n\n# 建立点分治 + 初始装载\ndecompose(1, -1)\nfor i in range(1, n+1):\n if is_red[i]:\n apply_update(i, +1)\n\nout_lines = []\nfor _ in range(q):\n t, v = ni(), ni()\n if t == 1:\n if is_red[v]:\n apply_update(v, -1)\n is_red[v] = 0\n else:\n apply_update(v, +1)\n is_red[v] = 1\n else:\n out_lines.append(str(query_sum(v)))\n\nsys.stdout.write("\n".join(out_lines))\n\n\n# Java \n\njava\nimport java.io.*;\nimport java.util.*;\n\n// 为了避免极端递归深度导致的栈溢出,主逻辑放到大栈线程里。\npublic class Main {\n static class FastScanner {\n private final InputStream in;\n private final byte[] buffer = new byte[1 << 16];\n private int ptr = 0, len = 0;\n FastScanner(InputStream is) { in = is; }\n private int read() throws IOException {\n if (ptr >= len) {\n len = in.read(buffer);\n ptr = 0;\n if (len <= 0) return -1;\n }\n return buffer[ptr++];\n }\n int nextInt() throws IOException {\n int c, sgn = 1, x = 0;\n do { c = read(); } while (c <= 32);\n if (c == '-') { sgn = -1; c = read(); }\n while (c > 32) {\n x = x * 10 + (c - '0');\n c = read();\n }\n return x * sgn;\n }\n }\n\n static class Edge { int to, w; Edge(int t,int w){this.to=t; this.w=w;} }\n\n static int n, q;\n static ArrayList<Edge>[] g;\n\n // 点分治\n static boolean[] removed;\n static int[] sz, cpar;\n\n // 每个点到各层重心的链(自上而下存,使用时倒序)\n static ArrayList<Integer>[] cpath;\n static ArrayList<Long>[] cdist;\n\n // 统计量\n static long[] totCnt, totDist;\n // 子方向统计:HashMap<下层重心, 值>\n @SuppressWarnings("unchecked")\n static HashMap<Integer, Long>[] subCnt, subDist;\n\n static int[] init;\n static boolean[] isRed;\n\n static void calcSize(int u, int p) {\n sz[u] = 1;\n for (Edge e : g[u]) {\n int v = e.to;\n if (v == p || removed[v]) continue;\n calcSize(v, u);\n sz[u] += sz[v];\n }\n }\n\n static int findCentroid(int u, int p, int tot) {\n for (Edge e : g[u]) {\n int v = e.to;\n if (v == p || removed[v]) continue;\n if (sz[v] > tot / 2) return findCentroid(v, u, tot);\n }\n return u;\n }\n\n static void collect(int u, int p, int cen, long dist) {\n cpath[u].add(cen);\n cdist[u].add(dist);\n for (Edge e : g[u]) {\n int v = e.to;\n if (v == p || removed[v]) continue;\n collect(v, u, cen, dist + e.w);\n }\n }\n\n static void decompose(int entry, int parent) {\n calcSize(entry, -1);\n int cen = findCentroid(entry, -1, sz[entry]);\n collect(cen, -1, cen, 0L);\n cpar[cen] = parent;\n removed[cen] = true;\n for (Edge e : g[cen]) {\n int v = e.to;\n if (!removed[v]) decompose(v, cen);\n }\n }\n\n static void applyUpdate(int u, int delta) {\n int prev = -1;\n int m = cpath[u].size();\n for (int i = m - 1; i >= 0; --i) {\n int c = cpath[u].get(i);\n long d = cdist[u].get(i);\n totCnt[c] += delta;\n totDist[c] += (long)delta * d;\n if (prev != -1) {\n subCnt[c].put(prev, subCnt[c].getOrDefault(prev, 0L) + delta);\n subDist[c].put(prev, subDist[c].getOrDefault(prev, 0L) + (long)delta * d);\n }\n prev = c;\n }\n }\n\n static long querySum(int u) {\n long ans = 0;\n int prev = -1;\n int m = cpath[u].size();\n for (int i = m - 1; i >= 0; --i) {\n int c = cpath[u].get(i);\n long d = cdist[u].get(i);\n ans += totDist[c] + totCnt[c] * d;\n if (prev != -1) {\n long sd = subDist[c].getOrDefault(prev, 0L);\n long sc = subCnt[c].getOrDefault(prev, 0L);\n ans -= sd + sc * d;\n }\n prev = c;\n }\n return ans;\n }\n\n public static void main(String[] args) throws Exception {\n new Thread(null, () -> {\n try {\n FastScanner fs = new FastScanner(System.in);\n n = fs.nextInt();\n q = fs.nextInt();\n init = new int[n+1];\n isRed = new boolean[n+1];\n\n g = new ArrayList[n+1];\n for (int i = 1; i <= n; ++i) g[i] = new ArrayList<>();\n for (int i = 1; i <= n; ++i) {\n init[i] = fs.nextInt();\n isRed[i] = init[i] == 1;\n }\n for (int i = 0; i < n-1; ++i) {\n int u = fs.nextInt(), v = fs.nextInt(), w = fs.nextInt();\n g[u].add(new Edge(v, w));\n g[v].add(new Edge(u, w));\n }\n\n removed = new boolean[n+1];\n sz = new int[n+1];\n cpar = new int[n+1];\n Arrays.fill(cpar, -1);\n\n cpath = new ArrayList[n+1];\n cdist = new ArrayList[n+1];\n for (int i = 1; i <= n; ++i) {\n cpath[i] = new ArrayList<>();\n cdist[i] = new ArrayList<>();\n }\n\n totCnt = new long[n+1];\n totDist = new long[n+1];\n\n subCnt = new HashMap[n+1];\n subDist = new HashMap[n+1];\n for (int i = 1; i <= n; ++i) {\n subCnt[i] = new HashMap<>();\n subDist[i] = new HashMap<>();\n }\n\n // 建点分治 + 初始装载集市状态\n decompose(1, -1);\n for (int i = 1; i <= n; ++i) if (isRed[i]) applyUpdate(i, +1);\n\n StringBuilder out = new StringBuilder();\n for (int i = 0; i < q; ++i) {\n int t = fs.nextInt(), v = fs.nextInt();\n if (t == 1) {\n if (isRed[v]) {\n applyUpdate(v, -1);\n isRed[v] = false;\n } else {\n applyUpdate(v, +1);\n isRed[v] = true;\n }\n } else {\n out.append(querySum(v)).append('\n');\n }\n }\n System.out.print(out.toString());\n } catch (Exception e) {\n e.printStackTrace();\n }\n }, "big-stack", 1 << 26).start(); // 提高线程栈\n }\n}\n
题目内容
在一片土地上,有 n 座城市,它们由 n−1 条道路连接,构成一棵树。城市编号为 1 到 n。最初,某些城市设有热闹的集市。每条道路具有正整数的长度,沿着道路通行所需的代价即为该长度。城市 u 与 v 之间的距离定义为二者之间唯一路径上所有道路的长度之和。
现在需要依次处理 q 个操作,操作分为两种类型:
- 切换状态:给定城市 x,将其集市状态翻转(原本开设则关闭,原本关闭则开设)。
- 查询距离和:给定城市 x,计算从 x 出发,到所有当前开设集市的城市距离之和。
请你对每个查询操作输出正确的答案。
数据规模与约定:城市数量 n 和操作次数 q 均满足 1≤n,q≤2imes105。道路长度均为正整数,且不超过 106。初始集市状态由 n 个 0 或 1 给出,1 表示该城市初始设有集市。保证至少有一个查询操作,且给定的道路构成一棵树。
请从“运行结果”或“历史提交”选择一条记录并点击「开始AI分析」
选择提交后点击「开始AI分析」