Number of Possible Sets of Closing Branches

hard graph floyd-warshall bitmask enumeration

Problem

There are n branches (0-indexed) connected by undirected weighted roads. You may close any subset of branches. A configuration is valid if, in the remaining graph, every pair of still-open branches can reach each other within maxDistance. Return the number of valid subsets of branches to keep open (the empty set counts).

With n small, enumerate all 2ⁿ subsets. For each subset run Floyd-Warshall on just the kept nodes and check that every open pair has shortest distance ≤ maxDistance.

Inputn = 3, maxDistance = 5, roads = [[0,1,2],[1,2,10],[0,2,10]]
Output5
Valid kept-sets: {}, {0}, {1}, {2}, {0,1}. {1,2}, {0,2}, {0,1,2} all have a pair beyond distance 5.

def number_of_sets(n, max_distance, roads):
    INF = float('inf')
    count = 0
    for mask in range(1 << n):
        d = [[INF] * n for _ in range(n)]
        for i in range(n):
            if mask & (1 << i):
                d[i][i] = 0
        for u, v, w in roads:
            if (mask & (1 << u)) and (mask & (1 << v)):
                d[u][v] = min(d[u][v], w)
                d[v][u] = min(d[v][u], w)
        for k in range(n):
            if not (mask & (1 << k)):
                continue
            for i in range(n):
                for j in range(n):
                    if d[i][k] + d[k][j] < d[i][j]:
                        d[i][j] = d[i][k] + d[k][j]
        ok = True
        for i in range(n):
            for j in range(i + 1, n):
                if (mask & (1 << i)) and (mask & (1 << j)) and d[i][j] > max_distance:
                    ok = False
        if ok:
            count += 1
    return count
function numberOfSets(n, maxDistance, roads) {
  const INF = Infinity;
  let count = 0;
  for (let mask = 0; mask < (1 << n); mask++) {
    const d = Array.from({ length: n }, () => new Array(n).fill(INF));
    for (let i = 0; i < n; i++) if (mask & (1 << i)) d[i][i] = 0;
    for (const [u, v, w] of roads) {
      if ((mask & (1 << u)) && (mask & (1 << v))) {
        d[u][v] = Math.min(d[u][v], w);
        d[v][u] = Math.min(d[v][u], w);
      }
    }
    for (let k = 0; k < n; k++) {
      if (!(mask & (1 << k))) continue;
      for (let i = 0; i < n; i++)
        for (let j = 0; j < n; j++)
          if (d[i][k] + d[k][j] < d[i][j]) d[i][j] = d[i][k] + d[k][j];
    }
    let ok = true;
    for (let i = 0; i < n; i++)
      for (let j = i + 1; j < n; j++)
        if ((mask & (1 << i)) && (mask & (1 << j)) && d[i][j] > maxDistance) ok = false;
    if (ok) count++;
  }
  return count;
}
class Solution {
    public int numberOfSets(int n, int maxDistance, int[][] roads) {
        int count = 0;
        final int INF = 1_000_000_000;
        for (int mask = 0; mask < (1 << n); mask++) {
            int[][] d = new int[n][n];
            for (int[] row : d) Arrays.fill(row, INF);
            for (int i = 0; i < n; i++) if ((mask & (1 << i)) != 0) d[i][i] = 0;
            for (int[] r : roads) {
                int u = r[0], v = r[1], w = r[2];
                if ((mask & (1 << u)) != 0 && (mask & (1 << v)) != 0) {
                    d[u][v] = Math.min(d[u][v], w);
                    d[v][u] = Math.min(d[v][u], w);
                }
            }
            for (int k = 0; k < n; k++) {
                if ((mask & (1 << k)) == 0) continue;
                for (int i = 0; i < n; i++)
                    for (int j = 0; j < n; j++)
                        if (d[i][k] + d[k][j] < d[i][j]) d[i][j] = d[i][k] + d[k][j];
            }
            boolean ok = true;
            for (int i = 0; i < n; i++)
                for (int j = i + 1; j < n; j++)
                    if ((mask & (1 << i)) != 0 && (mask & (1 << j)) != 0 && d[i][j] > maxDistance) ok = false;
            if (ok) count++;
        }
        return count;
    }
}
class Solution {
public:
    int numberOfSets(int n, int maxDistance, vector<vector<int>>& roads) {
        int count = 0;
        const int INF = 1e9;
        for (int mask = 0; mask < (1 << n); mask++) {
            vector<vector<int>> d(n, vector<int>(n, INF));
            for (int i = 0; i < n; i++) if (mask & (1 << i)) d[i][i] = 0;
            for (auto& r : roads) {
                int u = r[0], v = r[1], w = r[2];
                if ((mask & (1 << u)) && (mask & (1 << v))) {
                    d[u][v] = min(d[u][v], w);
                    d[v][u] = min(d[v][u], w);
                }
            }
            for (int k = 0; k < n; k++) {
                if (!(mask & (1 << k))) continue;
                for (int i = 0; i < n; i++)
                    for (int j = 0; j < n; j++)
                        if (d[i][k] + d[k][j] < d[i][j]) d[i][j] = d[i][k] + d[k][j];
            }
            bool ok = true;
            for (int i = 0; i < n; i++)
                for (int j = i + 1; j < n; j++)
                    if ((mask & (1 << i)) && (mask & (1 << j)) && d[i][j] > maxDistance) ok = false;
            if (ok) count++;
        }
        return count;
    }
};
Time: O(2ⁿ · n³) Space: O(n²)