Welcome to Subscribe On Youtube
3585. Find Weighted Median Node in Tree
Description
You are given an integer n and an undirected, weighted tree rooted at node 0 with n nodes numbered from 0 to n - 1. This is represented by a 2D array edges of length n - 1, where edges[i] = [ui, vi, wi] indicates an edge from node ui to vi with weight wi.
The weighted median node is defined as the first node x on the path from ui to vi such that the sum of edge weights from ui to x is greater than or equal to half of the total path weight.
You are given a 2D integer array queries. For each queries[j] = [uj, vj], determine the weighted median node along the path from uj to vj.
Return an array ans, where ans[j] is the node index of the weighted median for queries[j].
Example 1:
Input: n = 2, edges = [[0,1,7]], queries = [[1,0],[0,1]]
Output: [0,1]
Explanation:

| Query | Path | Edge Weights |
Total Path Weight |
Half | Explanation | Answer |
|---|---|---|---|---|---|---|
[1, 0] |
1 → 0 |
[7] |
7 | 3.5 | Sum from 1 → 0 = 7 >= 3.5, median is node 0. |
0 |
[0, 1] |
0 → 1 |
[7] |
7 | 3.5 | Sum from 0 → 1 = 7 >= 3.5, median is node 1. |
1 |
Example 2:
Input: n = 3, edges = [[0,1,2],[2,0,4]], queries = [[0,1],[2,0],[1,2]]
Output: [1,0,2]
Explanation:

| Query | Path | Edge Weights |
Total Path Weight |
Half | Explanation | Answer |
|---|---|---|---|---|---|---|
[0, 1] |
0 → 1 |
[2] |
2 | 1 | Sum from 0 → 1 = 2 >= 1, median is node 1. |
1 |
[2, 0] |
2 → 0 |
[4] |
4 | 2 | Sum from 2 → 0 = 4 >= 2, median is node 0. |
0 |
[1, 2] |
1 → 0 → 2 |
[2, 4] |
6 | 3 | Sum from 1 → 0 = 2 < 3.Sum from 1 → 2 = 2 + 4 = 6 >= 3, median is node 2. |
2 |
Example 3:
Input: n = 5, edges = [[0,1,2],[0,2,5],[1,3,1],[2,4,3]], queries = [[3,4],[1,2]]
Output: [2,2]
Explanation:

| Query | Path | Edge Weights |
Total Path Weight |
Half | Explanation | Answer |
|---|---|---|---|---|---|---|
[3, 4] |
3 → 1 → 0 → 2 → 4 |
[1, 2, 5, 3] |
11 | 5.5 | Sum from 3 → 1 = 1 < 5.5.Sum from 3 → 0 = 1 + 2 = 3 < 5.5.Sum from 3 → 2 = 1 + 2 + 5 = 8 >= 5.5, median is node 2. |
2 |
[1, 2] |
1 → 0 → 2 |
[2, 5] |
7 | 3.5 |
Sum from |
2 |
Constraints:
2 <= n <= 105edges.length == n - 1edges[i] == [ui, vi, wi]0 <= ui, vi < n1 <= wi <= 1091 <= queries.length <= 105queries[j] == [uj, vj]0 <= uj, vj < n- The input is generated such that
edgesrepresents a valid tree.
Solutions
Solution 1
-
class Solution { public int[] findMedian(int n, int[][] edges, int[][] queries) { int m = 32 - Integer.numberOfLeadingZeros(n); List<int[]>[] g = new List[n]; Arrays.setAll(g, i -> new ArrayList<>()); for (var e : edges) { int u = e[0], v = e[1], w = e[2]; g[u].add(new int[] {v, w}); g[v].add(new int[] {u, w}); } int[][] f = new int[n][m]; int[] p = new int[n]; int[] depth = new int[n]; long[] dist = new long[n]; Deque<Integer> q = new ArrayDeque<>(); q.offer(0); while (!q.isEmpty()) { int i = q.poll(); f[i][0] = p[i]; for (int j = 1; j < m; ++j) { f[i][j] = f[f[i][j - 1]][j - 1]; } for (var nxt : g[i]) { int j = nxt[0], w = nxt[1]; if (j != p[i]) { p[j] = i; depth[j] = depth[i] + 1; dist[j] = dist[i] + w; q.offer(j); } } } int[] ans = new int[queries.length]; for (int i = 0; i < queries.length; ++i) { int u = queries[i][0], v = queries[i][1]; if (u == v) { ans[i] = u; continue; } int x = u, y = v; if (depth[x] < depth[y]) { int t = x; x = y; y = t; } for (int j = m - 1; j >= 0; --j) { if (depth[x] - depth[y] >= (1 << j)) { x = f[x][j]; } } for (int j = m - 1; j >= 0; --j) { if (f[x][j] != f[y][j]) { x = f[x][j]; y = f[y][j]; } } if (x != y) { x = p[x]; } long w = dist[u] + dist[v] - 2 * dist[x]; if (2 * (dist[u] - dist[x]) >= w) { int cur = u; for (int j = m - 1; j >= 0; --j) { int k = f[cur][j]; if (depth[k] >= depth[x] && 2 * (dist[u] - dist[k]) < w) { cur = k; } } ans[i] = p[cur]; } else { int cur = v; for (int j = m - 1; j >= 0; --j) { int k = f[cur][j]; if (depth[k] > depth[x] && 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w) { cur = k; } } ans[i] = cur; } } return ans; } } -
class Solution { public: vector<int> findMedian(int n, vector<vector<int>>& edges, vector<vector<int>>& queries) { int m = 32 - __builtin_clz(n); vector<vector<pair<int, int>>> g(n); for (auto& e : edges) { int u = e[0], v = e[1], w = e[2]; g[u].emplace_back(v, w); g[v].emplace_back(u, w); } vector<vector<int>> f(n, vector<int>(m)); vector<int> p(n), depth(n); vector<long long> dist(n); queue<int> q; q.push(0); while (!q.empty()) { int i = q.front(); q.pop(); f[i][0] = p[i]; for (int j = 1; j < m; ++j) { f[i][j] = f[f[i][j - 1]][j - 1]; } for (auto [j, w] : g[i]) { if (j != p[i]) { p[j] = i; depth[j] = depth[i] + 1; dist[j] = dist[i] + w; q.push(j); } } } vector<int> ans; for (auto& qq : queries) { int u = qq[0], v = qq[1]; if (u == v) { ans.push_back(u); continue; } int x = u, y = v; if (depth[x] < depth[y]) { swap(x, y); } for (int j = m - 1; ~j; --j) { if (depth[x] - depth[y] >= (1 << j)) { x = f[x][j]; } } for (int j = m - 1; ~j; --j) { if (f[x][j] != f[y][j]) { x = f[x][j]; y = f[y][j]; } } if (x != y) { x = p[x]; } long long w = dist[u] + dist[v] - 2 * dist[x]; if (2 * (dist[u] - dist[x]) >= w) { int cur = u; for (int j = m - 1; ~j; --j) { int k = f[cur][j]; if (depth[k] >= depth[x] && 2 * (dist[u] - dist[k]) < w) { cur = k; } } ans.push_back(p[cur]); } else { int cur = v; for (int j = m - 1; ~j; --j) { int k = f[cur][j]; if (depth[k] > depth[x] && 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w) { cur = k; } } ans.push_back(cur); } } return ans; } }; -
class Solution: def findMedian( self, n: int, edges: List[List[int]], queries: List[List[int]] ) -> List[int]: m = n.bit_length() g = [[] for _ in range(n)] for u, v, w in edges: g[u].append((v, w)) g[v].append((u, w)) f = [[0] * m for _ in range(n)] p = [0] * n depth = [0] * n dist = [0] * n q = deque([0]) while q: i = q.popleft() f[i][0] = p[i] for j in range(1, m): f[i][j] = f[f[i][j - 1]][j - 1] for j, w in g[i]: if j != p[i]: p[j] = i depth[j] = depth[i] + 1 dist[j] = dist[i] + w q.append(j) ans = [] for u, v in queries: if u == v: ans.append(u) continue x, y = u, v if depth[x] < depth[y]: x, y = y, x for j in range(m - 1, -1, -1): if depth[x] - depth[y] >= (1 << j): x = f[x][j] for j in range(m - 1, -1, -1): if f[x][j] != f[y][j]: x, y = f[x][j], f[y][j] if x != y: x = p[x] w = dist[u] + dist[v] - 2 * dist[x] if 2 * (dist[u] - dist[x]) >= w: cur = u for j in range(m - 1, -1, -1): k = f[cur][j] if depth[k] >= depth[x] and 2 * (dist[u] - dist[k]) < w: cur = k ans.append(p[cur]) else: cur = v for j in range(m - 1, -1, -1): k = f[cur][j] if ( depth[k] > depth[x] and 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w ): cur = k ans.append(cur) return ans -
func findMedian(n int, edges [][]int, queries [][]int) []int { m := bits.Len(uint(n)) g := make([][][2]int, n) for _, e := range edges { u, v, w := e[0], e[1], e[2] g[u] = append(g[u], [2]int{v, w}) g[v] = append(g[v], [2]int{u, w}) } f := make([][]int, n) for i := range f { f[i] = make([]int, m) } p := make([]int, n) depth := make([]int, n) dist := make([]int, n) q := []int{0} for len(q) > 0 { i := q[0] q = q[1:] f[i][0] = p[i] for j := 1; j < m; j++ { f[i][j] = f[f[i][j-1]][j-1] } for _, nxt := range g[i] { j, w := nxt[0], nxt[1] if j != p[i] { p[j] = i depth[j] = depth[i] + 1 dist[j] = dist[i] + w q = append(q, j) } } } ans := make([]int, len(queries)) for i, qq := range queries { u, v := qq[0], qq[1] if u == v { ans[i] = u continue } x, y := u, v if depth[x] < depth[y] { x, y = y, x } for j := m - 1; j >= 0; j-- { if depth[x]-depth[y] >= 1<<j { x = f[x][j] } } for j := m - 1; j >= 0; j-- { if f[x][j] != f[y][j] { x, y = f[x][j], f[y][j] } } if x != y { x = p[x] } w := dist[u] + dist[v] - 2*dist[x] if 2*(dist[u]-dist[x]) >= w { cur := u for j := m - 1; j >= 0; j-- { k := f[cur][j] if depth[k] >= depth[x] && 2*(dist[u]-dist[k]) < w { cur = k } } ans[i] = p[cur] } else { cur := v for j := m - 1; j >= 0; j-- { k := f[cur][j] if depth[k] > depth[x] && 2*(dist[u]+dist[k]-2*dist[x]) >= w { cur = k } } ans[i] = cur } } return ans } -
function findMedian(n: number, edges: number[][], queries: number[][]): number[] { const m = 32 - Math.clz32(n); const g: number[][][] = Array.from({ length: n }, () => []); for (const [u, v, w] of edges) { g[u].push([v, w]); g[v].push([u, w]); } const f: number[][] = Array.from({ length: n }, () => Array(m).fill(0)); const p: number[] = Array(n).fill(0); const depth: number[] = Array(n).fill(0); const dist: number[] = Array(n).fill(0); const q: number[] = [0]; for (let qq = 0; qq < q.length; ++qq) { const i = q[qq]; f[i][0] = p[i]; for (let j = 1; j < m; ++j) { f[i][j] = f[f[i][j - 1]][j - 1]; } for (const [j, w] of g[i]) { if (j !== p[i]) { p[j] = i; depth[j] = depth[i] + 1; dist[j] = dist[i] + w; q.push(j); } } } const ans: number[] = []; for (const [u, v] of queries) { if (u === v) { ans.push(u); continue; } let x = u, y = v; if (depth[x] < depth[y]) { [x, y] = [y, x]; } for (let j = m - 1; j >= 0; --j) { if (depth[x] - depth[y] >= 1 << j) { x = f[x][j]; } } for (let j = m - 1; j >= 0; --j) { if (f[x][j] !== f[y][j]) { x = f[x][j]; y = f[y][j]; } } if (x !== y) { x = p[x]; } const w = dist[u] + dist[v] - 2 * dist[x]; if (2 * (dist[u] - dist[x]) >= w) { let cur = u; for (let j = m - 1; j >= 0; --j) { const k = f[cur][j]; if (depth[k] >= depth[x] && 2 * (dist[u] - dist[k]) < w) { cur = k; } } ans.push(p[cur]); } else { let cur = v; for (let j = m - 1; j >= 0; --j) { const k = f[cur][j]; if (depth[k] > depth[x] && 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w) { cur = k; } } ans.push(cur); } } return ans; } -
public class Solution { public int[] FindMedian(int n, int[][] edges, int[][] queries) { int m = 32 - BitOperations.LeadingZeroCount((uint)n); List<int[]>[] g = new List<int[]>[n]; for (int i = 0; i < n; ++i) { g[i] = new List<int[]>(); } foreach (var e in edges) { int u = e[0], v = e[1], w = e[2]; g[u].Add(new int[] { v, w }); g[v].Add(new int[] { u, w }); } int[][] f = new int[n][]; for (int i = 0; i < n; ++i) { f[i] = new int[m]; } int[] p = new int[n]; int[] depth = new int[n]; long[] dist = new long[n]; Queue<int> q = new Queue<int>(); q.Enqueue(0); while (q.Count > 0) { int i = q.Dequeue(); f[i][0] = p[i]; for (int j = 1; j < m; ++j) { f[i][j] = f[f[i][j - 1]][j - 1]; } foreach (var nxt in g[i]) { int j = nxt[0], w = nxt[1]; if (j != p[i]) { p[j] = i; depth[j] = depth[i] + 1; dist[j] = dist[i] + w; q.Enqueue(j); } } } int[] ans = new int[queries.Length]; for (int i = 0; i < queries.Length; ++i) { int u = queries[i][0], v = queries[i][1]; if (u == v) { ans[i] = u; continue; } int x = u, y = v; if (depth[x] < depth[y]) { int t = x; x = y; y = t; } for (int j = m - 1; j >= 0; --j) { if (depth[x] - depth[y] >= (1 << j)) { x = f[x][j]; } } for (int j = m - 1; j >= 0; --j) { if (f[x][j] != f[y][j]) { x = f[x][j]; y = f[y][j]; } } if (x != y) { x = p[x]; } long w = dist[u] + dist[v] - 2 * dist[x]; if (2 * (dist[u] - dist[x]) >= w) { int cur = u; for (int j = m - 1; j >= 0; --j) { int k = f[cur][j]; if (depth[k] >= depth[x] && 2 * (dist[u] - dist[k]) < w) { cur = k; } } ans[i] = p[cur]; } else { int cur = v; for (int j = m - 1; j >= 0; --j) { int k = f[cur][j]; if (depth[k] > depth[x] && 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w) { cur = k; } } ans[i] = cur; } } return ans; } } -
use std::collections::VecDeque; impl Solution { pub fn find_median(n: i32, edges: Vec<Vec<i32>>, queries: Vec<Vec<i32>>) -> Vec<i32> { let n = n as usize; let m = 32 - (n as u32).leading_zeros() as usize; let mut g = vec![vec![]; n]; for e in &edges { let u = e[0] as usize; let v = e[1] as usize; let w = e[2] as i64; g[u].push((v, w)); g[v].push((u, w)); } let mut f = vec![vec![0; m]; n]; let mut p = vec![0; n]; let mut depth = vec![0; n]; let mut dist = vec![0i64; n]; let mut q = VecDeque::new(); q.push_back(0); while let Some(i) = q.pop_front() { f[i][0] = p[i]; for j in 1..m { f[i][j] = f[f[i][j - 1]][j - 1]; } for &(j, w) in &g[i] { if j != p[i] { p[j] = i; depth[j] = depth[i] + 1; dist[j] = dist[i] + w; q.push_back(j); } } } let mut ans = Vec::with_capacity(queries.len()); for qq in &queries { let u = qq[0] as usize; let v = qq[1] as usize; if u == v { ans.push(u as i32); continue; } let (mut x, mut y) = (u, v); if depth[x] < depth[y] { std::mem::swap(&mut x, &mut y); } for j in (0..m).rev() { if depth[x] - depth[y] >= (1 << j) { x = f[x][j]; } } for j in (0..m).rev() { if f[x][j] != f[y][j] { x = f[x][j]; y = f[y][j]; } } if x != y { x = p[x]; } let w = dist[u] + dist[v] - 2 * dist[x]; if 2 * (dist[u] - dist[x]) >= w { let mut cur = u; for j in (0..m).rev() { let k = f[cur][j]; if depth[k] >= depth[x] && 2 * (dist[u] - dist[k]) < w { cur = k; } } ans.push(p[cur] as i32); } else { let mut cur = v; for j in (0..m).rev() { let k = f[cur][j]; if depth[k] > depth[x] && 2 * (dist[u] + dist[k] - 2 * dist[x]) >= w { cur = k; } } ans.push(cur as i32); } } ans } }