A星路径搜索:原理与Demo实现

A* 路径搜索:原理与 Demo


目录

  1. 为什么需要 A*

  2. A* 算法是什么

  3. 启发函数与可接受性

  4. Demo 实战

  5. 常见踩坑


本文讲清 A* 路径搜索原理,并用 Python + Qt Demo 直观看到 Open/Closed 扩展过程。Demo 同时提供 Dijkstra(h≡0)对照,帮助理解 A* 中 h 的作用。

项 说明
环境 Ubuntu 22.04 · Python 3.10+
依赖 pip3 install PyQt5
运行 python3 astar_demo.py
源码 同目录 astar_demo.py、README.md

1. 为什么需要 A*

路径搜索:在栅格地图上,从起点 S 到终点 G 找一条 代价最小且可通行 的路径。

算法 核心 最优性 效率
BFS 按层扩展 单位代价下最优 慢
Dijkstra 按 g 扩展 最优 较慢
贪心 Best-First 只看 h 不保证 快但易绕路
A* f = g + h h 可接受时最优 通常最快

Dijkstra 只信 g;A* 用 g + h,让搜索更朝终点收敛。


2. A* 算法是什么

A* 是 启发式图搜索:每次从 Open List 取出 f 最小 的节点扩展,直到到达终点。

2.1 核心公式

f(n) = g(n) + h(n)

符号 含义
g(n) 起点 → n 的 实际代价
h(n) n → 终点的 启发估计
f(n) 总评估,用于排序

2.2 Open / Closed List

列表 作用
Open 待扩展,按 f 升序(常用最小堆)
Closed 已扩展,防重复(常用集合)

2.3 搜索流程

1. 起点入 Open,g=0
2. 取 f 最小节点 current
3. current == goal → 沿 parent 回溯路径
4. current 入 Closed,扩展邻居:
   更新更优 g,计算 f=g+h,入 Open
5. Open 空 → 无解

2.4 与 Dijkstra

h(n) ≡ 0 时,A* 退化为 Dijkstra(只按 g 排序)。Demo 用 heuristic=lambda _: 0 即可切换对照。


3. 启发函数与可接受性

3.1 常用 h(四邻接栅格用曼哈顿)

名称 公式
曼哈顿
欧几里得 √((x₁−x₂)² + (y₁−y₂)²)
对角线 max(

Demo 四邻接移动,默认 曼哈顿距离:

h(node) = abs(node.row - goal.row) + abs(node.col - goal.col)

3.2 可接受性(Admissible)

若 h(n) 永不高估 真实最短代价,则 A* 保证最优。

h 情况 结果
可接受(不高估) 最优路径
高估 可能更快但不保证最优
h ≡ 0 等价 Dijkstra

4. Demo 实战

4.1 安装与运行

python3 --version          # ≥ 3.10
pip3 install PyQt5         # 或 sudo apt install python3-pyqt5
python3 astar_demo.py

4.2 A* 核心实现

astar_demo.py 中搜索入口为 astar(map_data, heuristic=None),默认 h 为曼哈顿距离:

def manhattan(a, b):
    return abs(a[0] - b[0]) + abs(a[1] - b[1])
​
def astar(map_data, heuristic=None, on_expand=None):
    h = heuristic or (lambda node: manhattan(node, goal))
    open_heap = [(h(start), 0, 0, start)]
    g_score, parent, closed = {start: 0}, {start: None}, set()
​
    while open_heap:
        _, g, _, current = heapq.heappop(open_heap)
        if current in closed or g > g_score[current]:
            continue
        closed.add(current)
        if current == goal:
            return 沿 parent 回溯
​
        for neighbor in 四邻接:
            tentative_g = g + 1
            if tentative_g < g_score.get(neighbor, inf):
                g_score[neighbor] = tentative_g
                parent[neighbor] = current
                f = tentative_g + h(neighbor)   # A* 关键行
                heapq.heappush(open_heap, (f, tentative_g, counter, neighbor))

A* 行为:优先扩展 f 最小节点 → Closed 区域 朝终点收窄;浅蓝=Closed,黄=Open,橙=当前,红线=路径。

4.3 与 Dijkstra 对照(可选)

astar(map_data, heuristic=lambda _: 0)   # h=0,按 g 均匀扩散

同一地图下路径步数通常相同,但 A* 扩展节点更少:

对比 A* Dijkstra
排序键 f = g + h f = g
扩展形状 朝终点收窄 均匀扩散
扩展节点 更少 更多

回字地图实测:路径 72 步;A* 扩展 388 节点,Dijkstra 418 节点。

复现:「加载回字地图」→ A*「立即完成」→ 重置 → Dijkstra「立即完成」,对比状态栏扩展数与蓝色区域。

4.4 GUI 速查

控件 作用
生成迷宫 / 加载回字地图 切换地图
算法下拉框 A* ↔ Dijkstra
开始搜索 / 立即完成 动画 / 直接出路径
重置 清空搜索状态

完整按钮说明与常见问题见 README.md;完整源码见 astar_demo.py。


5. 常见踩坑

现象 处理
GUI 起不来 pip3 install PyQt5
回字地图格子错乱 用最新 astar_demo.py(行宽已统一)
路径与文档不一致 多条最优路径时步数可能相同,以状态栏实测为准
配图不显示 png 需与 md 同目录

6. 完整代码

#!/usr/bin/env python3

from __future__ import annotations

import heapq
import random
import sys
from dataclasses import dataclass, field
from typing import Callable, Iterable

# ---------------------------------------------------------------------------
# 地图与 A* 核心
# ---------------------------------------------------------------------------

NEIGHBORS = ((-1, 0), (1, 0), (0, -1), (0, 1))

# 所有行等宽,避免锯齿网格导致切换地图后绘制/搜索异常
DEFAULT_MAZE = [
    "################################################################",
    "#S.............................................................#",
    "#######.#######.#######.#######.#######.#######.#######.######.#",
    "#..............................................................#",
    "#######.#######.#######.#######.#######.#######.#######.######.#",
    "#..............................................................#",
    "#######.#######.#######.#######.#######.#######.#######.######.#",
    "#..............................................................#",
    "#######.#######.#######.#######.#######.#######.#######.######.#",
    "#..............................................................#",
    "#######.#######.#######.#######.#######.#######.#######.######.#",
    "#.............................................................G#",
    "################################################################",
]


@dataclass
class SearchResult:
    path: list[tuple[int, int]] | None
    closed: set[tuple[int, int]]
    expansion_order: list[tuple[int, int]]
    open_remaining: set[tuple[int, int]]
    nodes_expanded: int
    max_open_size: int
    g_score: dict[tuple[int, int], int]


@dataclass
class MapData:
    grid: list[list[int]]
    start: tuple[int, int]
    goal: tuple[int, int]

    @property
    def rows(self) -> int:
        return len(self.grid)

    @property
    def cols(self) -> int:
        return len(self.grid[0])


def parse_maze(lines: Iterable[str]) -> MapData:
    grid: list[list[int]] = []
    start = goal = None
    width = None
    for r, line in enumerate(lines):
        if width is None:
            width = len(line)
        elif len(line) != width:
            raise ValueError(f"地图第 {r} 行长度 {len(line)} 与首行 {width} 不一致")
        row: list[int] = []
        for c, ch in enumerate(line):
            if ch == "#":
                row.append(1)
            elif ch == "S":
                row.append(0)
                start = (r, c)
            elif ch == "G":
                row.append(0)
                goal = (r, c)
            else:
                row.append(0)
        grid.append(row)
    if start is None or goal is None:
        raise ValueError("地图必须包含 S(起点)和 G(终点)")
    return MapData(grid, start, goal)


def generate_maze(rows: int, cols: int, seed: int | None = None) -> MapData:

    if rows < 5 or cols < 5:
        raise ValueError("迷宫尺寸至少 5x5")
    if rows % 2 == 0:
        rows -= 1
    if cols % 2 == 0:
        cols -= 1

    rng = random.Random(seed)
    grid = [[1] * cols for _ in range(rows)]
    stack = [(1, 1)]
    grid[1][1] = 0

    while stack:
        r, c = stack[-1]
        directions = [(0, 2), (0, -2), (2, 0), (-2, 0)]
        rng.shuffle(directions)
        carved = False
        for dr, dc in directions:
            nr, nc = r + dr, c + dc
            if 0 < nr < rows - 1 and 0 < nc < cols - 1 and grid[nr][nc] == 1:
                grid[r + dr // 2][c + dc // 2] = 0
                grid[nr][nc] = 0
                stack.append((nr, nc))
                carved = True
                break
        if not carved:
            stack.pop()

    start = (1, 1)
    goal = (rows - 2, cols - 2)
    grid[start[0]][start[1]] = 0
    grid[goal[0]][goal[1]] = 0
    return MapData(grid, start, goal)


def manhattan(a: tuple[int, int], b: tuple[int, int]) -> int:
    return abs(a[0] - b[0]) + abs(a[1] - b[1])


def astar(
    map_data: MapData,
    heuristic: Callable[[tuple[int, int]], int] | None = None,
    on_expand: Callable[[tuple[int, int], set[tuple[int, int]], set[tuple[int, int]]], None] | None = None,
) -> SearchResult:
    grid = map_data.grid
    start, goal = map_data.start, map_data.goal
    rows, cols = map_data.rows, map_data.cols
    h = heuristic or (lambda node: manhattan(node, goal))

    open_heap: list[tuple[int, int, int, tuple[int, int]]] = []
    counter = 0
    g_score: dict[tuple[int, int], int] = {start: 0}
    parent: dict[tuple[int, int], tuple[int, int] | None] = {start: None}
    heapq.heappush(open_heap, (h(start), 0, counter, start))
    counter += 1

    closed: set[tuple[int, int]] = set()
    open_set = {start}
    expansion_order: list[tuple[int, int]] = []
    max_open_size = 1

    while open_heap:
        max_open_size = max(max_open_size, len(open_set))
        _, g, _, current = heapq.heappop(open_heap)
        if current in closed:
            continue
        if g > g_score.get(current, float("inf")):
            continue

        open_set.discard(current)
        closed.add(current)
        expansion_order.append(current)

        if on_expand:
            on_expand(current, closed, open_set)

        if current == goal:
            return SearchResult(
                path=_reconstruct(parent, goal),
                closed=closed,
                expansion_order=expansion_order,
                open_remaining=set(open_set),
                nodes_expanded=len(closed),
                max_open_size=max_open_size,
                g_score=dict(g_score),
            )

        for dx, dy in NEIGHBORS:
            nx, ny = current[0] + dx, current[1] + dy
            neighbor = (nx, ny)
            if not (0 <= nx < rows and 0 <= ny < cols):
                continue
            if grid[nx][ny] == 1:
                continue

            tentative_g = g + 1
            if tentative_g >= g_score.get(neighbor, float("inf")):
                continue

            g_score[neighbor] = tentative_g
            parent[neighbor] = current
            f = tentative_g + h(neighbor)
            heapq.heappush(open_heap, (f, tentative_g, counter, neighbor))
            counter += 1
            open_set.add(neighbor)

    return SearchResult(
        path=None,
        closed=closed,
        expansion_order=expansion_order,
        open_remaining=set(open_set),
        nodes_expanded=len(closed),
        max_open_size=max_open_size,
        g_score=dict(g_score),
    )


def _reconstruct(parent: dict[tuple[int, int], tuple[int, int] | None], goal: tuple[int, int]) -> list[tuple[int, int]]:
    path: list[tuple[int, int]] = []
    cur: tuple[int, int] | None = goal
    while cur is not None:
        path.append(cur)
        cur = parent.get(cur)
    path.reverse()
    return path


# ---------------------------------------------------------------------------
# Qt 图形界面
# ---------------------------------------------------------------------------

@dataclass
class ViewState:
    closed: set[tuple[int, int]] = field(default_factory=set)
    open_set: set[tuple[int, int]] = field(default_factory=set)
    current: tuple[int, int] | None = None
    path: list[tuple[int, int]] | None = None
    show_path: bool = False


def run_gui() -> None:
    try:
        from PyQt5.QtCore import QPointF, Qt, QTimer
        from PyQt5.QtGui import QColor, QFont, QPainter, QPainterPath, QPen
        from PyQt5.QtWidgets import (
            QApplication,
            QComboBox,
            QHBoxLayout,
            QLabel,
            QMainWindow,
            QPushButton,
            QSlider,
            QSpinBox,
            QVBoxLayout,
            QWidget,
        )
    except ImportError as exc:
        print("GUI 模式需要 PyQt5,请执行: pip install PyQt5", file=sys.stderr)
        raise SystemExit(1) from exc

    class MazeCanvas(QWidget):
        CELL = 14
        MARGIN = 12
        COLORS = {
            "bg": QColor("#1e1e2e"),
            "wall": QColor("#45475a"),
            "free": QColor("#cdd6f4"),
            "closed": QColor("#89b4fa"),
            "open": QColor("#f9e2af"),
            "current": QColor("#fab387"),
            "path": QColor("#f38ba8"),
            "path_line": QColor("#e64553"),
            "start": QColor("#a6e3a1"),
            "goal": QColor("#eba0ac"),
            "grid": QColor("#313244"),
        }

        def __init__(self, parent=None):
            super().__init__(parent)
            self._map: MapData | None = None
            self._view = ViewState()
            self.setMinimumSize(640, 480)

        def set_map(self, map_data: MapData) -> None:
            self._map = map_data
            self._view = ViewState()
            self.update()
            self.repaint()

        def set_view(self, view: ViewState) -> None:
            self._view = view
            self.update()

        def paintEvent(self, event):  # noqa: N802
            painter = QPainter(self)
            painter.setRenderHint(QPainter.Antialiasing)
            painter.fillRect(self.rect(), self.COLORS["bg"])

            if not self._map:
                painter.setPen(QColor("#cdd6f4"))
                painter.drawText(self.rect(), Qt.AlignCenter, "点击「生成迷宫」或「加载回字地图」开始")
                return

            map_data = self._map
            view = self._view
            path_set = set(view.path or [])
            cell, margin = self.CELL, self.MARGIN
            map_w, map_h = map_data.cols * cell, map_data.rows * cell
            scale = min((self.width() - margin * 2) / map_w, (self.height() - margin * 2) / map_h, 2.5)
            scale = max(scale, 0.4)
            ox = (self.width() - map_w * scale) / 2
            oy = (self.height() - map_h * scale) / 2

            def rect_for(r, c):
                x = ox + c * cell * scale
                y = oy + r * cell * scale
                s = cell * scale
                return x, y, s

            def center(r, c):
                x, y, s = rect_for(r, c)
                return QPointF(x + s / 2, y + s / 2)

            for r in range(map_data.rows):
                for c in range(map_data.cols):
                    pos = (r, c)
                    x, y, s = rect_for(r, c)
                    if map_data.grid[r][c] == 1:
                        color = self.COLORS["wall"]
                    elif pos in path_set and view.show_path:
                        color = self.COLORS["path"]
                    elif pos == view.current:
                        color = self.COLORS["current"]
                    elif pos in view.closed:
                        color = self.COLORS["closed"]
                    elif pos in view.open_set:
                        color = self.COLORS["open"]
                    elif pos == map_data.start:
                        color = self.COLORS["start"]
                    elif pos == map_data.goal:
                        color = self.COLORS["goal"]
                    else:
                        color = self.COLORS["free"]
                    painter.fillRect(int(x), int(y), int(s) + 1, int(s) + 1, color)

            if view.show_path and view.path and len(view.path) > 1:
                pen = QPen(self.COLORS["path_line"], max(2.0, 3 * scale))
                pen.setCapStyle(Qt.RoundCap)
                pen.setJoinStyle(Qt.RoundJoin)
                painter.setPen(pen)
                painter.setBrush(Qt.NoBrush)
                path_obj = QPainterPath(center(*view.path[0]))
                for node in view.path[1:]:
                    path_obj.lineTo(center(*node))
                painter.drawPath(path_obj)
                painter.setBrush(self.COLORS["path_line"])
                painter.setPen(Qt.NoPen)
                radius = max(2.0, 2.5 * scale)
                for node in view.path:
                    pt = center(*node)
                    painter.drawEllipse(pt, radius, radius)

            font = QFont("Sans", max(7, int(7 * scale)))
            font.setBold(True)
            painter.setFont(font)
            for pos, label in ((map_data.start, "S"), (map_data.goal, "G")):
                pt = center(*pos)
                painter.setPen(QColor("#1e1e2e"))
                painter.drawText(int(pt.x() - 6 * scale), int(pt.y() + 4 * scale), label)

    class MainWindow(QMainWindow):
        def __init__(self):
            super().__init__()
            self.setWindowTitle("A* 路径搜索可视化")
            self.resize(1100, 720)
            self._map: MapData | None = None
            self._result: SearchResult | None = None
            self._open_history: list[set[tuple[int, int]]] = []
            self._step = 0
            self._animating = False
            self._maze_rows = 31
            self._maze_cols = 51
            self._maze_seed = 42
            self._using_hui_map = False
            self._timer = QTimer(self)
            self._timer.timeout.connect(self._on_tick)
            self._canvas = MazeCanvas()
            self._status = QLabel("就绪")
            self._status.setStyleSheet("color: #cdd6f4; padding: 4px;")
            self._rows_spin = QSpinBox()
            self._rows_spin.setRange(11, 61)
            self._rows_spin.setSingleStep(2)
            self._rows_spin.setValue(31)
            self._cols_spin = QSpinBox()
            self._cols_spin.setRange(11, 81)
            self._cols_spin.setSingleStep(2)
            self._cols_spin.setValue(51)
            self._seed_spin = QSpinBox()
            self._seed_spin.setRange(0, 99999)
            self._seed_spin.setValue(42)
            self._speed = QSlider(Qt.Horizontal)
            self._speed.setRange(1, 50)
            self._speed.setValue(15)
            self._speed_label = QLabel("15 ms/步")
            self._algo = QComboBox()
            self._algo.addItems(["A* (曼哈顿)", "Dijkstra (h=0)"])
            btn_style = (
                "QPushButton { background: #45475a; color: #cdd6f4; padding: 6px 12px; border-radius: 4px; }"
                "QPushButton:hover { background: #585b70; }"
            )
            btn_gen = QPushButton("生成迷宫")
            btn_gen.clicked.connect(self._generate_maze)
            btn_default = QPushButton("加载回字地图")
            btn_default.clicked.connect(self._load_default)
            btn_run = QPushButton("开始搜索")
            btn_run.clicked.connect(self._start_search)
            btn_instant = QPushButton("立即完成")
            btn_instant.clicked.connect(self._finish_instant)
            btn_reset = QPushButton("重置")
            btn_reset.clicked.connect(self._reset_view)
            for btn in (btn_gen, btn_default, btn_run, btn_instant, btn_reset):
                btn.setStyleSheet(btn_style)
            ctrl = QWidget()
            layout = QVBoxLayout(ctrl)
            row1 = QHBoxLayout()
            for w in (QLabel("行:"), self._rows_spin, QLabel("列:"), self._cols_spin,
                      QLabel("种子:"), self._seed_spin, self._algo, btn_gen, btn_default):
                row1.addWidget(w)
            layout.addLayout(row1)
            row2 = QHBoxLayout()
            for w in (btn_run, btn_instant, btn_reset, QLabel("动画速度:"), self._speed, self._speed_label):
                row2.addWidget(w)
            layout.addLayout(row2)
            layout.addWidget(self._status)
            for w in (self._rows_spin, self._cols_spin, self._seed_spin, self._algo):
                w.setStyleSheet("color: #cdd6f4; background: #313244;")
            central = QWidget()
            main_layout = QVBoxLayout(central)
            main_layout.addWidget(self._canvas, stretch=1)
            main_layout.addWidget(ctrl)
            central.setStyleSheet("background: #181825;")
            self.setCentralWidget(central)
            self._speed.valueChanged.connect(self._on_speed_changed)
            self._generate_maze()

        def _on_speed_changed(self, value: int) -> None:
            self._speed_label.setText(f"{value} ms/步")
            if self._animating:
                self._timer.setInterval(value)

        def _heuristic(self):
            return None if self._algo.currentIndex() == 0 else (lambda _: 0)

        def _apply_map(self, map_data: MapData, label: str) -> None:
            """切换地图并清空搜索状态,强制画布刷新。"""
            self._timer.stop()
            self._animating = False
            self._result = None
            self._open_history = []
            self._step = 0
            self._map = map_data
            self._canvas.set_map(map_data)
            self._status.setText(label)

        def _generate_maze(self) -> None:
            # 从回字地图切回时,恢复进入回字前的迷宫参数
            if self._using_hui_map:
                self._rows_spin.setValue(self._maze_rows)
                self._cols_spin.setValue(self._maze_cols)
                self._seed_spin.setValue(self._maze_seed)

            seed = self._seed_spin.value() or None
            m = generate_maze(self._rows_spin.value(), self._cols_spin.value(), seed)
            self._using_hui_map = False
            self._maze_rows, self._maze_cols = m.rows, m.cols
            self._maze_seed = self._seed_spin.value()
            self._rows_spin.blockSignals(True)
            self._cols_spin.blockSignals(True)
            self._rows_spin.setValue(m.rows)
            self._cols_spin.setValue(m.cols)
            self._rows_spin.blockSignals(False)
            self._cols_spin.blockSignals(False)
            self._apply_map(
                m,
                f"迷宫 {m.rows}×{m.cols},起点 {m.start},终点 {m.goal},"
                f"曼哈顿估计 {manhattan(m.start, m.goal)} 步",
            )

        def _load_default(self) -> None:
            if not self._using_hui_map:
                self._maze_rows = self._rows_spin.value()
                self._maze_cols = self._cols_spin.value()
                self._maze_seed = self._seed_spin.value()
            m = parse_maze(DEFAULT_MAZE)
            self._using_hui_map = True
            self._rows_spin.blockSignals(True)
            self._cols_spin.blockSignals(True)
            self._rows_spin.setValue(m.rows)
            self._cols_spin.setValue(m.cols)
            self._rows_spin.blockSignals(False)
            self._cols_spin.blockSignals(False)
            self._apply_map(m, f"回字地图 {m.rows}×{m.cols},起点 {m.start},终点 {m.goal}")

        def _start_search(self) -> None:
            if not self._map:
                return
            self._timer.stop()
            self._open_history = []

            def on_expand(_current, _closed, open_set):
                self._open_history.append(set(open_set))

            self._result = astar(self._map, heuristic=self._heuristic(), on_expand=on_expand)
            self._step = 0
            self._animating = True
            self._timer.start(self._speed.value())

        def _finish_instant(self) -> None:
            if not self._map:
                return
            if not self._result:
                self._result = astar(self._map, heuristic=self._heuristic())
            self._timer.stop()
            self._animating = False
            self._show_final()
            self._update_status(done=True)

        def _reset_view(self) -> None:
            self._timer.stop()
            self._animating = False
            self._step = 0
            self._result = None
            self._open_history = []
            if self._map:
                self._canvas.set_view(ViewState())
                self._status.setText("已重置,可重新搜索")

        def _on_tick(self) -> None:
            if not self._result or not self._map:
                self._timer.stop()
                return
            order = self._result.expansion_order
            if self._step >= len(order):
                self._timer.stop()
                self._animating = False
                self._show_final()
                self._update_status(done=True)
                return
            closed = set(order[: self._step + 1])
            current = order[self._step]
            open_set = self._open_history[self._step] if self._step < len(self._open_history) else set()
            self._canvas.set_view(ViewState(closed=closed, open_set=open_set, current=current))
            self._status.setText(f"搜索中… 已扩展 {self._step + 1}/{len(order)} 节点,当前 {current}")
            self._step += 1

        def _show_final(self) -> None:
            if not self._result or not self._map:
                return
            self._canvas.set_view(ViewState(
                closed=self._result.closed,
                open_set=self._result.open_remaining,
                path=self._result.path,
                show_path=True,
            ))

        def _update_status(self, done: bool = False) -> None:
            if not self._result or not self._map:
                return
            r, algo = self._result, self._algo.currentText()
            if r.path:
                msg = f"[{algo}] 路径步数 {len(r.path) - 1},扩展 {r.nodes_expanded} 节点,Open 峰值 {r.max_open_size}"
            else:
                msg = f"[{algo}] 无可行路径,已扩展 {r.nodes_expanded} 节点"
            if done:
                msg += " — 轨迹已绘制"
            self._status.setText(msg)

    app = QApplication(sys.argv)
    app.setStyle("Fusion")
    win = MainWindow()
    win.show()
    sys.exit(app.exec_())


# ---------------------------------------------------------------------------
# 入口
# ---------------------------------------------------------------------------

def main() -> None:
    run_gui()


if __name__ == "__main__":
    main()


很好的教程,很棒的思路,使我只需要一个python和Qt依赖就可以可视化地很直观地学习与理解导航算法 :+1::nerd_face: :+1: