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 1 → 0 = 2 < 3.5.
Sum from 1 → 2 = 2 + 5 = 7 >= 3.5, median is node 2.

2

 

Constraints:

  • 2 <= n <= 105
  • edges.length == n - 1
  • edges[i] == [ui, vi, wi]
  • 0 <= ui, vi < n
  • 1 <= wi <= 109
  • 1 <= queries.length <= 105
  • queries[j] == [uj, vj]
  • 0 <= uj, vj < n
  • The input is generated such that edges represents 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
        }
    }
    
    

All Problems

All Solutions