How to efficiently count all concatenations of 2-tuples into longer chains in Python

Viewed 97

Let us say that we would like to build a long (metal) chain which will be composed of smaller links, chained together. I know what the length of the chain should be: n. The links are represented as 2-tuples: (a, b). We may chain links together if and only if they share the same element at the side by which they would be chained.
I am given a list of lists of length n-1 - links - which represents all links available to me at each position of the chain. For example:

links = [
    [
        ('a', 1),
        ('a', 2),
        ('a', 3),
        ('b', 1),
    ],
    [
        (1, 'A'),
        (2, 'A'),
        (2, 'B'),
    ],
    [
        ('A', 'a'),
        ('B', 'a'),
        ('B', 'b'),
    ]
]

In this case the length of the final chain will be: n = 4.
Here we may generate these possible chains:

('a', 1, 'A', 'a')
('b', 1, 'A', 'a')
('a', 2, 'A', 'a')
('a', 2, 'B', 'a')
('a', 2, 'B', 'b')

This procedure is quite similar to forming a long line with domino puzzles, however I cannot rotate the tiles.

My task is that given such an input list I need to calculate all possible distinct chains of length n that may be created. The case above is a simplified toy example but in reality the chain's length may be as high as 1000 and I may be able to use tens of different links at each specific position. However, I know that for sure for each link available at position i there exists another link at position i-1 which is compatible to it.

I wrote a very naive solution with iterates over all links from beginning to end and merges them together, growing all possible versions of the final chain:


    # THIS CODE WAS ORIGINALLY BUGGED ONCE I POSTED IT
    # BUT IS FIXED NOW

    # initiate chains with links that could make up
    # the first position, then: iteratively grow them
    chains = links[0]
    
    # seach for all possible paths:
    # iterate over all positions
    for position in links[1:]:
        
        # temp array to help me grow the chain
        temp = []
            
        # over each chain in the current set of chains
        for chain in chains:

            # over each link in a given position
            for link in position:
                
                # check if the chain and link are chainable
                if chain[-1] == link[0]:
                    
                    # append new link to a pre-existing chain
                    temp.append(chain + tuple([link[1]]))
        
        # overwrite the current list of chains
        chains = temp

This solution works fine, i.e. I am quite convinced it returns a correct result. However, it is extremely slow, I need to speed it up, preferably ~100x. Therefore I think I need to employ a smart algorithm to count all the possibilities, not a brute-force concatenation as above... Since I only need to count the chains, not enumerate them, maybe there would be a backtracking procedure which would start from each possible final link and multiply possibilities along the way; in the end adding up over all final links? I have some vague ideas but cannot really nail this down...

2 Answers

Since counting is enough, let's just do that, and then it takes a split second for large cases as well.

from collections import defaultdict, Counter

def count_chains(links):
    chains = defaultdict(lambda: 1)
    for position in links:
        temp = Counter()
        for a, b in position:
            temp[b] += chains[a]
        chains = temp
    return sum(chains.values())

It does pretty much the same as yours, except instead of chains being a list of chains ending in some b-values, I'm using a Counter of chains ending in some b-values: chains[b] tells me how many chains end in b. And Counters (and defaultdict) are dictionaries, so I don't have to search and check for matching connectors, I just look them up.

The backwards compatibility means we might better go backwards, so we're not tracking dead ends, but I don't think it would help much if at all (depends on your data).

For example for links = [[(1, 1), (1, 2), (2, 1), (2, 2)]] * 1000, it takes about 2 ms to compute the number of chains, which is:

21430172143725346418968500981200036211228096234110672148875007767407021022498722449863967576313917162551893458351062936503742905713846280871969155149397149607869135549648461970842149210124742283755908364306092949967163882534797535118331087892154125829142392955373084335320859663305248773674411336138752

Try it online!

Here is my solution with Graph data structure approach which will be more efficient that yours which has cubic time complexity O(n3).

import timeit

class Node:
    def __init__(self, val, next=None):
        self.val = val
        self.next = next if next else []
        self.visited = False
        self.mem = None
    def __repr__(self):
        return f'<{self.val} {self.mem}>'
        
def epsi_sol(links, len_):
    nodes = {}
    start_nodes = set(i[0] for i in links[0]) # {'a', 'b'}

    # constructiong the graph
    for i in links:
        for j in i:
            if j[0] not in nodes:
                nodes[j[0]] = Node(j[0])
            if j[1] not in nodes:
                nodes[j[1]] = Node(j[1])

            nodes[j[0]].next.append(nodes[j[1]])

    def find_chain_with_length(node, length, valid_length):
        if length +1 == valid_length:
            return 1

        # if already visited just return
        if node.visited:
            return 0 
        
        if node.mem is not None:
            return node.mem

        # if this is not leaf node
        # we will mark it visited
        node.visited = True
        temp_count = 0
        for each_neighbor in node.next:
            temp_count += find_chain_with_length(each_neighbor, length+1, valid_length)
        # after visiting mark it unvisited
        node.visited = False
        node.mem = temp_count
        return temp_count

    solution_count = 0
    for each_start_node in start_nodes:
        solution_count += find_chain_with_length(nodes[each_start_node],0, len_)
    return solution_count

from collections import defaultdict, Counter

def kelly_sol(links, len_):
    chains = defaultdict(lambda: 1)
    for position in links[:len_]:
        temp = Counter()
        for a, b in position:
            temp[b] += chains[a]
        chains = temp
    return sum(chains.values())

def mac_sol(links, len_):
    chains = links[0]
    
    # seach for all possible paths:
    # iterate over all positions
    for position in links[1:]:
        
        # temp array to help me grow the chain
        temp = []
            
        # over each chain in the current set of chains
        for chain in chains:

            # over each link in a given position
            for link in position:
                
                # check if the chain and link are chainable
                if chain[-1] == link[0]:
                    
                    # append new link to a pre-existing chain
                    temp.append(chain + tuple([link[1]]))
        
        # overwrite the current list of chains
        chains = temp
    return len(chains)

# tests
for n in range(100, 1000, 100):
    links = [[(f'1_{i}', f'1_{i+1}'), (f'1_{i}', f'2_{i+1}'), (f'2_{i}', f'1_{i+1}'), (f'2_{i}', f'2_{i+1}')] for i in range(n)]
    print('-'*50)
    print(f'kelly_sol({n}) => {timeit.timeit(lambda: kelly_sol(links, n+1), number=2)} seconds')
    print(f'epsi_sol({n}) => {timeit.timeit(lambda: epsi_sol(links, n+1), number=2)} seconds')
    if n <50:
        print(f'mac_sol({n}) => {timeit.timeit(lambda: mac_sol(links, n+1), number=2)} seconds')
    print('-'*50)
--------------------------------------------------
kelly_sol(100) => 0.0013752000000000209 seconds
epsi_sol(100) => 0.0022567999999996147 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(200) => 0.0026332999999993945 seconds
epsi_sol(200) => 0.004522500000000207 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(300) => 0.003924899999999454 seconds
epsi_sol(300) => 0.006861199999999457 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(400) => 0.005278099999999952 seconds
epsi_sol(400) => 0.012999699999999947 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(500) => 0.006728900000000593 seconds
epsi_sol(500) => 0.01406989999999908 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(600) => 0.00828249999999997 seconds
epsi_sol(600) => 0.015398799999999824 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(700) => 0.009703200000000578 seconds
epsi_sol(700) => 0.01597070000000045 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(800) => 0.009961999999999804 seconds
epsi_sol(800) => 0.0196051999999991 seconds
--------------------------------------------------
--------------------------------------------------
kelly_sol(900) => 0.014800799999999725 seconds
epsi_sol(900) => 0.02183789999999952 seconds
--------------------------------------------------
Related