Python实现A*搜索运行极慢,近距起止节点也耗时久该如何优化?
代码性能问题核心原因
1. 最核心的Bug:Node类缺少相等性判断逻辑
你在用neighbour in openlist、neighbour in closed做包含判断时,Python默认比对对象的内存地址而非节点位置。你每次生成新的next = Node(current,[x,y])对象时,哪怕这个位置的节点已经在openlist/closed中,新生成的对象和已有对象地址不同,判断结果永远为False,会重复往openlist里塞入大量相同位置的节点,导致openlist规模指数级膨胀,哪怕很小的地图也会跑很久。
2. open列表取最小f值的效率太低
你每次遍历整个openlist找最小f值,时间复杂度是O(n),n是openlist的长度,地图越大速度越慢。标准实现应该用优先队列(最小堆),取最小元素的时间复杂度是O(logn)。
3. 包含判断效率低
你直接在列表上做in判断,时间复杂度是O(n),本身就很慢,建议用集合/字典存储已经访问过的节点位置,查询时间复杂度为O(1)。
4. 其他小问题
- 代码里用到了
math.sqrt但没有导入math模块 - 欧氏距离的开平方可以省略,比较f值大小时,h的平方和原值的大小顺序完全一致,可减少计算量
优化修改方案
首先修改Node类,增加__eq__方法支持按位置判断相等,增加__hash__方法方便存入集合:
import math import heapq class Node(object): def __init__(self, parent = None, position = None): self.position = position self.parent = parent self.g = 0 self.h = 0 self.f = 0 def __eq__(self, other): return self.position == other.position def __hash__(self): return hash(tuple(self.position)) def euclidean_distance(self, x_end, y_end): return math.sqrt(abs(self.position[0]-x_end)**2 + abs(self.position[1]-y_end)**2) def pos(self): return self.position
再修改astar函数的实现,改用优先队列,用集合存储已关闭的位置:
maze = [[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0], [0, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0], [1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]] def astar(): open_heap = [] closed_set = set() open_pos_set = set() # 额外存储open里的位置,避免重复插入 start = Node(None,[0,0]) destination = [4,8] # 堆里存( f值, 插入计数, 节点 ),加计数是为了避免f值相同时比较节点报错 heapq.heappush(open_heap, (start.f, 0, start)) open_pos_set.add(tuple(start.position)) count = 1 adjacent = [[1,0],[0,1],[-1,0],[0,-1]] while open_heap: current_f, _, current = heapq.heappop(open_heap) current_pos_tuple = tuple(current.position) if current_pos_tuple in closed_set: continue closed_set.add(current_pos_tuple) if current.pos() == destination: print("ok") path = [] current_node = current while current_node is not None: path.append(current_node.position) current_node = current_node.parent return path for candidate in adjacent: x = candidate[0]+current.pos()[0] y = candidate[1]+current.pos()[1] if x > (len(maze) - 1) or x < 0 or y > (len(maze[-1])-1) or y < 0: continue if maze[x][y] != 0: continue next_pos = [x,y] next_pos_tuple = tuple(next_pos) if next_pos_tuple in closed_set: continue next_node = Node(current, next_pos) next_node.g = current.g + 1 next_node.h = math.sqrt((next_pos[0]-destination[0])**2 + (next_pos[1]-destination[1])**2) next_node.f = next_node.g + next_node.h if next_pos_tuple not in open_pos_set or next_node.g < current.g: heapq.heappush(open_heap, (next_node.f, count, next_node)) open_pos_set.add(next_pos_tuple) count +=1
优化后的代码运行当前测试地图可毫秒级返回结果。
内容的提问来源于stack exchange,提问作者vikram bhamre
相关产品推荐
相关产品推荐

