Is there a way to find all unions within a nested list so that the result is a nested list in which no list has a common element?

Viewed 41

I've got a function that returns a nested list such as [[1,2], [2,3], [5,6]]. To find all unions inside the nested list I've tried using sets unions but I'm not sure how to filter down in a loop as the size of the nested list is not constant.
Is there a way through list comprehenision or nested for loops?

Examples:

Given the input [[1,2], [2,3], [5,6]] --> the output would be [[1,2,3], [5,6]]

[[0], [2, 5, 6], [1, 3, 5], [2, 4], [3, 5], [1, 2, 4], [1]] --> [[0], [1,2,3,4,5,6]]

1 Answers

You can use the disjoint sets data structure. The data structure supports two operations:

  • Union, where we take two members in the data structure and assign those elements (as well as all the other members in both of their respective sets) to the same disjoint set (in essence, unioning the two sets that the members belong to).
  • Find, which efficiently looks up which disjoint set an element belongs to. We use this at the end to read off the groupings into a list of lists.

The runtime of this is (almost) linear in the total number of elements in the ragged list.

Here is a pure Python implementation:

class DisjointSets:
    def __init__(self, n):
        self.elements = [-1] * n

    def union(self, first, second):
        first_root, second_root = self.find(first), self.find(second)
        if first_root == second_root:
            return False
        elif first_root < second_root:
            self.elements[first_root] += self.elements[second_root]
            self.elements[second_root] = first_root
            return True
        else:
            self.elements[second_root] += self.elements[first_root]
            self.elements[first_root] = second_root
            return True

    def find(self, target):
        if self.elements[target] < 0:
            return target
        self.elements[target] = self.find(self.elements[target])
        return self.elements[target]
        
data = [[0], [2, 5, 6], [1, 3, 5], [2, 4], [3, 5], [1, 2, 4], [1]] 

# Map elements in the data list to indices
# in the list of the disjoint sets data structure.
item_indices = {}
idx = 0
for sublist in data:
    for item in sublist:
        if item not in item_indices:
            item_indices[item] = idx
            idx += 1

# Bind every element in a sublist to a disjoint set.
# If the element has been seen previously, the data structure
# will correctly union the sets across the sublists.
disjoint_sets = DisjointSets(len(item_indices))
for sublist in data:
    for item1, item2 in zip(sublist, sublist[1:]):
        disjoint_sets.union(item_indices[item1], item_indices[item2])

# Read off the resulting groups.
result = {}
for item in item_indices:
    group_id = disjoint_sets.find(item_indices[item])
    if group_id not in result:
        result[group_id] = []
    result[group_id].append(item)

print(list(result.values()))

This code is a bit long, but the upside is that it requires no dependencies. If you want something more concise, networkx has an implementation of the disjoint sets data structure that you can use.

Related