Python: implementação do algoritmo Astar

Sep 02 2020

Eu implementei o algoritmo Astar para um problema em um juiz online relacionado às posições inicial e final do labirinto, juntamente com uma grade que representa o labirinto. Eu imprimo o comprimento do caminho junto com o caminho em si. A seguir está a implementação em Python usando a distância euclidiana:

import heapq, math, sys

infinity = float('inf')

class AStar():

    def __init__(self, start, grid, height, width):
        self.start, self.grid, self.height, self.width = start, grid, height, width

    class Node():
        def __init__(self, position, fscore=infinity, gscore=infinity, parent = None):
            self.fscore, self.gscore, self.position, self.parent = fscore, gscore, position, parent
            
        def __lt__(self, comparator):
            return self.fscore < comparator.fscore

    def heuristic(self, end, distance = "Euclidean"):
        (x1, y1), (x2, y2) = self.start, end
        if (distance == "Manhattan"):
            return abs(x1 - x2) + abs(y1 - y2)
        return math.sqrt((x2 - x1)**2 + (y2 - y1)**2)

    def nodeNeighbours(self, pos):
        (x, y) = pos
        return [(dx, dy) for (dx, dy) in [(x + 1, y), (x - 1, y), (x, y + 1), (x, y - 1)] if 0 <= dx < self.width and 0 <= dy < self.height and self.grid[dy][dx] == 0]

    def getPath(self, endPoint):
        current, path = endPoint, []
        while current.position != self.start:
            path.append(current.position)
            current = current.parent
        path.append(self.start)
        return list(reversed(path))

    def computePath(self, end):
        openList, closedList, nodeDict = [], [], {}
        currentNode = AStar.Node(self.start, fscore=self.heuristic(end), gscore = 0)
        heapq.heappush(openList, currentNode)
        while openList:
            currentNode = heapq.heappop(openList)
            if currentNode.position == end:
                return self.getPath(currentNode)
            else:
                closedList.append(currentNode)
                neighbours = []
                for toCheck in self.nodeNeighbours(currentNode.position):
                    if toCheck not in nodeDict.keys():
                        nodeDict[toCheck] = AStar.Node(toCheck)
                        neighbours.append(nodeDict[toCheck])
                
                for neighbour in neighbours:
                    newGscore = currentNode.gscore + 1
                    if neighbour in openList and newGscore < neighbour.gscore:
                        openList.remove(neighbour)
                    if newGscore < neighbour.gscore and neighbour in closedList:
                        closedList.remove(neighbour)
                    if neighbour not in openList and neighbour not in closedList:
                        neighbour.gscore = newGscore
                        neighbour.fscore = neighbour.gscore + self.heuristic(neighbour.position)
                        neighbour.parent = currentNode
                        heapq.heappush(openList, neighbour)
                    heapq.heapify(openList)
        return None
        
if __name__ == '__main__':
    
    sys.stdin = open('input.txt', 'r')
    sys.stdout = open('output.txt', 'w')
    
    matrix = [[int(num) for num in line.split()] for line in sys.stdin]
    size = matrix.pop(0)
    coordinates = matrix.pop(0)
    n, m = size[0], size[1]
    x1, y1, y2, x2 = coordinates[0], coordinates[1], coordinates[2], coordinates[3]
    path = AStar((x1-1, y1-1), matrix, n, m).computePath((y2-1, x2-1))
    print(len(path))
    for pos in path:
        print(pos[0] + 1, pos[1] + 1)

Respostas

5 Carcigenicate Sep 02 2020 at 21:24
self.start, self.grid, self.height, self.width = start, grid, height, width

Eu não colocaria todos na mesma linha dessa forma. Acho que seria muito mais fácil ler em várias linhas:

self.start = start
self.grid = grid
self.height = height
self.width = width

Eu provavelmente teria a Nodeclasse de nível superior em vez de aninhada. Não acho que você está ganhando muito por tê-lo dentro AStar. Você poderia nomeá-lo _Nodepara torná-lo "privado do módulo", de modo que tentar importá-lo para outro arquivo potencialmente gerará avisos.

Em Node's __lt__implementação, eu não chamaria o segundo parâmetro comparator. Um comparador é algo que compara, enquanto, neste caso, é apenas outro nó. other_nodeou algo seria mais apropriado.


Em heuristic, eu pessoalmente faria uso de um elselá:

if (distance == "Manhattan"):
    return abs((x1 - x2) + abs(y1 - y2))
else:
    return math.sqrt((x2 - x1)**2 + (y2 - y1)**2)

Deixa claro que apenas uma das linhas será executada. Pessoalmente, só desprezo o elseem um caso como esse se iffor uma verificação de pré-condição de "saída antecipada" e quero evitar aninhar todo o resto da função dentro de um bloco. Isso não é um problema aqui.


nodeNeighbors( que deveria sernode_neighbors ) seria mais claro dividido em várias linhas:

def nodeNeighbours(self, pos):
    (x, y) = pos
    return [(dx, dy)
            for (dx, dy) in [(x + 1, y), (x - 1, y), (x, y + 1), (x, y - 1)]
            if 0 <= dx < self.width and 0 <= dy < self.height and self.grid[dy][dx] == 0]

Acho que isso torna muito mais fácil ver o que está acontecendo nele.


Novamente, em muitos lugares você está atribuindo duas ou mais variáveis ​​em uma linha:

(x1, y1), (x2, y2) = self.start, end
current, path = endPoint, []
openList, closedList, nodeDict = [], [], {}
x1, y1, y2, x2 = coordinates[0], coordinates[1], coordinates[2], coordinates[3]

Eu iria quebrar isso. Especialmente quando você chega a 3+ em uma linha, para que o leitor veja qual variável corresponde a qual valor, ele precisará contar a partir da esquerda em vez de apenas verificar o que está em cada lado de a =.


Em computePath, parece que closedListdeveria ser um conjunto. Não parece que a ordem importe com ele e neighbour in closedListserá mais rápido com um conjunto do que com uma lista. Parece que openListé necessário ser uma lista devido ao fato de ter sido passada para heapify.


Eu não acho que eu reatribuiria stdine stdout. A reatribuição de stdinparece completamente desnecessária e a alteração stdouttornará mais difícil depurar posteriormente usando printinstruções. Você não deseja necessariamente que todo o texto impresso seja enviado para o arquivo.

Se necessário, você pode especificar em qual arquivo deseja imprimir ao imprimir:

with open('output.txt', 'w') as out_f:
    print("To file!", file=out_f)