Find Critical and Pseudo Critical Edges in Minimum Spanning Tree
The drill: A connected, weighted graph can have several different minimum spanning trees, all tied at the same total weight. Sort every edge into critical (in ALL of them), pseudo-critical (in SOME of them), or neither.
A connected, weighted graph arrives, and it can have more than one minimum spanning tree — different edge sets that all tie at the same lowest total weight.
Every edge in the graph gets sorted into one of three groups: critical, meaning it shows up in every minimum spanning tree without exception; pseudo-critical, meaning it shows up in at least one but not all of them; or neither, meaning no minimum spanning tree ever needs it.
The output lists the critical edges and the pseudo-critical edges separately, each identified by its position in the original edge list rather than by its endpoints.
- graph stays connected and modest in size, up to roughly a hundred edges
- edge weights are positive and can repeat across different edges
- edges are identified by their index in the input list
- output is two index lists: critical edges, then pseudo-critical edges
HINT 1 THE NUDGE
Every minimum spanning tree ties at the same total weight, even though the edge sets can differ. That single baseline weight is the yardstick both categories get measured against.
HINT 2 THE STRUCTURE
Test each edge two separate ways against that baseline: what happens to the achievable weight if this edge is banned outright, and separately, what's the best achievable weight if this edge is forced in before anything else runs?
HINT 3 ONE STEP FROM THE ANSWER
Ban edge i and rebuild the tree — a strictly higher weight (or a disconnected graph) makes it critical. Otherwise, force edge i in first and rebuild — landing back on the baseline weight makes it pseudo-critical.
Triangle, 3 distinct weights: (0,1)=1, (1,2)=2, (0,2)=3. First find the baseline MST weight with Kruskal, then test each edge's necessity.
class Solution:
def findCriticalAndPseudoCriticalEdges(self, n: int, edges: List[List[int]]) -> List[List[int]]:
base = self._weight(n, edges, -1, -1)
critical, pseudo = [], []
for i in range(len(edges)):
if self._weight(n, edges, i, -1) > base:
critical.append(i)
elif self._weight(n, edges, -1, i) == base:
pseudo.append(i)
return [critical, pseudo]
def _find(self, parent, x):
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
def _weight(self, n, edges, exclude, include):
order = sorted(range(len(edges)), key=lambda i: edges[i][2])
parent = list(range(n))
total = 0
count = 0
if include != -1:
u, v, w = edges[include]
parent[self._find(parent, u)] = self._find(parent, v)
total += w
count += 1
for idx in order:
if idx == exclude or idx == include:
continue
if count == n - 1:
break
u, v, w = edges[idx]
pu, pv = self._find(parent, u), self._find(parent, v)
if pu != pv:
parent[pu] = pv
total += w
count += 1
return total if count == n - 1 else float('inf')class Solution:
def findCriticalAndPseudoCriticalEdges(self, n: int, edges: List[List[int]]) -> List[List[int]]:
base = self._weight(n, edges, -1, -1)
critical, pseudo = [], []
for i in range(len(edges)):
if self._weight(n, edges, i, -1) > base:
critical.append(i)
elif self._weight(n, edges, -1, i) == base:
pseudo.append(i)
return [critical, pseudo]
def _weight(self, n, edges, exclude, include):
order = sorted(range(len(edges)), key=lambda i: edges[i][2])
adj = collections.defaultdict(list)
total = 0
count = 0
if include != -1:
u, v, w = edges[include]
adj[u].append(v)
adj[v].append(u)
total += w
count += 1
for idx in order:
if idx == exclude or idx == include:
continue
if count == n - 1:
break
u, v, w = edges[idx]
if not self._connected(adj, u, v):
adj[u].append(v)
adj[v].append(u)
total += w
count += 1
return total if count == n - 1 else float('inf')
def _connected(self, adj, u, v):
if u == v:
return True
visited = {u}
stack = [u]
while stack:
node = stack.pop()
if node == v:
return True
for nb in adj[node]:
if nb not in visited:
visited.add(nb)
stack.append(nb)
return Falseclass Solution {
public int[][] findCriticalAndPseudoCriticalEdges(int n, int[][] edges) {
int base = weight(n, edges, -1, -1);
List<Integer> critical = new ArrayList<>();
List<Integer> pseudo = new ArrayList<>();
for (int i = 0; i < edges.length; i++) {
if (weight(n, edges, i, -1) > base) {
critical.add(i);
} else if (weight(n, edges, -1, i) == base) {
pseudo.add(i);
}
}
return new int[][] { toArray(critical), toArray(pseudo) };
}
private int find(int[] parent, int x) {
while (parent[x] != x) {
parent[x] = parent[parent[x]];
x = parent[x];
}
return x;
}
private int weight(int n, int[][] edges, int exclude, int include) {
Integer[] order = new Integer[edges.length];
for (int i = 0; i < edges.length; i++) order[i] = i;
Arrays.sort(order, (a, b) -> edges[a][2] - edges[b][2]);
int[] parent = new int[n];
for (int i = 0; i < n; i++) parent[i] = i;
int total = 0, count = 0;
if (include != -1) {
int u = edges[include][0], v = edges[include][1], w = edges[include][2];
parent[find(parent, u)] = find(parent, v);
total += w;
count++;
}
for (int idx : order) {
if (idx == exclude || idx == include) continue;
if (count == n - 1) break;
int u = edges[idx][0], v = edges[idx][1], w = edges[idx][2];
int pu = find(parent, u), pv = find(parent, v);
if (pu != pv) {
parent[pu] = pv;
total += w;
count++;
}
}
return count == n - 1 ? total : Integer.MAX_VALUE;
}
private int[] toArray(List<Integer> list) {
int[] a = new int[list.size()];
for (int i = 0; i < a.length; i++) a[i] = list.get(i);
return a;
}
}class Solution {
public int[][] findCriticalAndPseudoCriticalEdges(int n, int[][] edges) {
int base = weight(n, edges, -1, -1);
List<Integer> critical = new ArrayList<>();
List<Integer> pseudo = new ArrayList<>();
for (int i = 0; i < edges.length; i++) {
if (weight(n, edges, i, -1) > base) {
critical.add(i);
} else if (weight(n, edges, -1, i) == base) {
pseudo.add(i);
}
}
return new int[][] { toArray(critical), toArray(pseudo) };
}
private int weight(int n, int[][] edges, int exclude, int include) {
Integer[] order = new Integer[edges.length];
for (int i = 0; i < edges.length; i++) order[i] = i;
Arrays.sort(order, (a, b) -> edges[a][2] - edges[b][2]);
Map<Integer, List<Integer>> adj = new HashMap<>();
int total = 0, count = 0;
if (include != -1) {
int u = edges[include][0], v = edges[include][1], w = edges[include][2];
adj.computeIfAbsent(u, x -> new ArrayList<>()).add(v);
adj.computeIfAbsent(v, x -> new ArrayList<>()).add(u);
total += w;
count++;
}
for (int idx : order) {
if (idx == exclude || idx == include) continue;
if (count == n - 1) break;
int u = edges[idx][0], v = edges[idx][1], w = edges[idx][2];
if (!connected(adj, u, v)) {
adj.computeIfAbsent(u, x -> new ArrayList<>()).add(v);
adj.computeIfAbsent(v, x -> new ArrayList<>()).add(u);
total += w;
count++;
}
}
return count == n - 1 ? total : Integer.MAX_VALUE;
}
private boolean connected(Map<Integer, List<Integer>> adj, int u, int v) {
if (u == v) return true;
Set<Integer> visited = new HashSet<>();
visited.add(u);
Deque<Integer> stack = new ArrayDeque<>();
stack.push(u);
while (!stack.isEmpty()) {
int node = stack.pop();
if (node == v) return true;
for (int nb : adj.getOrDefault(node, Collections.emptyList())) {
if (!visited.contains(nb)) {
visited.add(nb);
stack.push(nb);
}
}
}
return false;
}
private int[] toArray(List<Integer> list) {
int[] a = new int[list.size()];
for (int i = 0; i < a.length; i++) a[i] = list.get(i);
return a;
}
}✓ CHIP-TIMED — ALL 4 SOLUTIONS RAN GREEN AGAINST SELF-AUTHORED CASES IN CI · JDK 21 · CPYTHON 3.12 · NOTHING PUBLISHES RED