DeepSeek LeetCode 3786. 树组的交互代价总和 Rust实现

📅 2026/7/31 4:17:39
DeepSeek    LeetCode 3786. 树组的交互代价总和 Rust实现
问题描述给定一棵 n 个节点的无向树节点编号 0 到 n-1数组 group[i] 表示节点 i 所属的分组。两个节点 u 和 v 的交互代价为树上它们之间唯一路径的边数。要求返回所有同组无序节点对的交互代价总和。---核心思路边贡献统计法直接枚举所有同组节点对并计算路径长度会达到 O(n²)不可行。关键转化总代价 每条边被同组节点对的路径经过的次数之和。对于任意一条边将其从树中删除会把树分成两部分。假设某一分组在这条边一侧子树中有 x 个节点该组全局总数为 k则该组中路径经过这条边的节点对数量为 x * (k - x)。因此只需一次 DFS 遍历统计每个子树中各分组的节点数然后累加每条边的贡献即可。---Rust 实现递归版rustuse std::collections::HashMap;impl Solution {pub fn interaction_costs(n: i32, edges: VecVeci32, group: Veci32) - i64 {let n n as usize;// 1. 构建邻接表let mut adj vec![Vec::new(); n];for e in edges {let u e[0] as usize;let v e[1] as usize;adj[u].push(v);adj[v].push(u);}// 2. 离散化分组标签因为分组标签可能不连续let mut group_map HashMap::new();for g in group {group_map.entry(g).or_insert(0);}let m group_map.len(); // 不同分组的数量// 为每个分组分配紧凑索引let mut idx 0;for (_, v) in group_map.iter_mut() {*v idx;idx 1;}// 将原始分组标签转换为紧凑索引let gid: Vecusize group.iter().map(|g| *group_map.get(g).unwrap()).collect();// 3. 统计全局各分组节点总数let mut total vec![0; m];for id in gid {total[id] 1;}// 4. cnt[u][g] 以 u 为根的子树中分组 g 的节点数let mut cnt vec![vec![0; m]; n];let mut ans 0i64;// DFS 递归函数fn dfs(u: usize,parent: usize,adj: VecVecusize,gid: Vecusize,total: Veci32,cnt: mut VecVeci32,ans: mut i64,) {cnt[u][gid[u]] 1; // 当前节点自身for v in adj[u] {if v parent { continue; }dfs(v, u, adj, gid, total, cnt, ans);// 计算边 (u, v) 对答案的贡献for g in 0..total.len() {if total[g] 2 { continue; } // 该组少于2个节点无贡献let in_subtree cnt[v][g] as i64; // 子树中该组节点数let out_subtree total[g] as i64 - in_subtree; // 子树外该组节点数if in_subtree 0 out_subtree 0 {*ans in_subtree * out_subtree;}}// 合并子树的统计信息到当前节点for g in 0..total.len() {cnt[u][g] cnt[v][g];}}}dfs(0, n, adj, gid, total, mut cnt, mut ans);ans}}---Rust 实现迭代版避免栈溢出rustuse std::collections::HashMap;impl Solution {pub fn interaction_costs(n: i32, edges: VecVeci32, group: Veci32) - i64 {let n n as usize;// 1. 构建邻接表let mut adj vec![Vec::new(); n];for e in edges {let u e[0] as usize;let v e[1] as usize;adj[u].push(v);adj[v].push(u);}// 2. 离散化分组标签let mut group_map HashMap::new();for g in group {group_map.entry(g).or_insert(0);}let m group_map.len();let mut idx 0;for (_, v) in group_map.iter_mut() {*v idx;idx 1;}let gid: Vecusize group.iter().map(|g| *group_map.get(g).unwrap()).collect();// 3. 统计全局各分组节点总数let mut total vec![0; m];for id in gid {total[id] 1;}// 4. 迭代 DFS 获取遍历顺序let mut parent vec![n; n]; // n 作为哨兵值let mut order Vec::with_capacity(n);let mut stack vec![0];parent[0] n; // 根节点的父节点标记为 nwhile let Some(u) stack.pop() {order.push(u);for v in adj[u] {if v parent[u] { continue; }parent[v] u;stack.push(v);}}// 5. 逆序遍历从叶子到根累计贡献let mut cnt vec![vec![0; m]; n];let mut ans 0i64;for u in order.iter().rev() {cnt[u][gid[u]] 1; // 当前节点自身for v in adj[u] {if v parent[u] { continue; } // 只处理子节点// 计算边 (u, v) 对答案的贡献for g in 0..m {if total[g] 2 { continue; }let in_subtree cnt[v][g] as i64;let out_subtree total[g] as i64 - in_subtree;if in_subtree 0 out_subtree 0 {ans in_subtree * out_subtree;}}// 合并子节点计数for g in 0..m {cnt[u][g] cnt[v][g];}}}ans}}---分组标签范围固定时的简化版本如果题目保证分组标签范围为 1~20与 LeetCode 3786 原题一致rustimpl Solution {pub fn interaction_costs(n: i32, edges: VecVeci32, group: Veci32) - i64 {const MAX_GROUP: usize 20;let n n as usize;// 构建邻接表let mut adj vec![Vec::new(); n];for e in edges {let u e[0] as usize;let v e[1] as usize;adj[u].push(v);adj[v].push(u);}// 统计全局各分组节点总数let mut total vec![0; MAX_GROUP 1];for g in group {total[g as usize] 1;}// 迭代 DFS 获取遍历顺序let mut parent vec![n; n];let mut order Vec::with_capacity(n);let mut stack vec![0];parent[0] n;while let Some(u) stack.pop() {order.push(u);for v in adj[u] {if v parent[u] { continue; }parent[v] u;stack.push(v);}}// 逆序遍历累计贡献let mut cnt vec![vec![0; MAX_GROUP 1]; n];let mut ans 0i64;for u in order.iter().rev() {cnt[u][group[u] as usize] 1;for v in adj[u] {if v parent[u] { continue; }for g in 1..MAX_GROUP {if total[g] 2 { continue; }let in_subtree cnt[v][g] as i64;let out_subtree total[g] as i64 - in_subtree;if in_subtree 0 out_subtree 0 {ans in_subtree * out_subtree;}}for g in 1..MAX_GROUP {cnt[u][g] cnt[v][g];}}}ans}}---代码说明1. 离散化处理Rust 中无法直接用不连续的分组标签作为数组索引因此使用 HashMap 进行离散化将原始标签映射到 0..m-1 的紧凑索引。2. 两种 DFS 实现· 递归版代码简洁但 Rust 默认栈较小深度过大可能栈溢出。· 迭代版使用显式栈避免递归适合大规模数据推荐使用。3. 核心计算对于边 (u, v)cnt[v][g] 为子树中该组节点数total[g] - cnt[v][g] 为子树外同组节点数。乘积即为该组中路径经过这条边的节点对数量。4. 复杂度分析· 时间复杂度O(n · m)其中 m 为不同分组的数量≤ 20。· 空间复杂度O(n · m) 用于存储 cnt 数组加上 O(n) 的邻接表。---测试示例rustfn main() {let n 4;let edges vec![vec![0,1], vec![0,2], vec![2,3]];let group vec![1, 2, 1, 2];let result Solution::interaction_costs(n, edges, group);println!({}, result); // 输出: 3}解释同组节点对 (0,2) 路径长度为 1(1,3) 路径长度为 2总和为 3。---注意事项· 答案可能很大使用 i64 存储结果。· Rust 递归深度限制默认较小n 较大时请使用迭代版。· 若使用固定分组范围版本需要确认题目中分组标签确实在 1~20 范围内。