每日算法 — 使用java实现迷宫探索:Q-Learning强化学习与状态空间最优策略

引言

在人工智能领域,强化学习(Reinforcement Learning)是一种让智能体通过与环境交互来自主学习最优决策策略的强大范式。与监督学习不同,强化学习不需要标注数据,智能体通过不断试错、获得奖励反馈,逐步优化自身行为。Q-Learning作为最经典的免模型(Model-Free)强化学习算法之一,因其简洁高效而被广泛应用于游戏AI、机器人控制、路径规划等场景。

迷宫探索是理解强化学习的绝佳切入点:环境简单直观,状态空间可控,且目标明确——让智能体从起点出发,通过学习找到通往出口的最短路径。本文将用Java从零实现一个基于Q-Learning的迷宫探索系统,深入讲解状态空间建模、Q-Table更新机制、ε-贪心策略等核心概念,并提供可直接运行的完整代码。

Q-Learning算法核心原理

马尔可夫决策过程

迷宫探索问题可以形式化为一个马尔可夫决策过程(MDP),由以下四元组定义:

  • 状态空间 S:智能体在迷宫中的所有可能位置
  • 动作空间 A:每个状态下可执行的动作(上、下、左、右)
  • 转移函数 P(s’|s,a):执行动作a后从状态s转移到s’的概率(在确定性迷宫中为1或0)
  • 奖励函数 R(s,a,s’):执行动作后获得的即时奖励

Q-Table与贝尔曼方程

Q-Learning的核心是维护一张Q-Table,记录每个状态-动作对的价值:

Q(s, a) = 在状态s执行动作a后,遵循最优策略所能获得的累积折扣奖励期望值

Q值的更新遵循贝尔曼最优方程

Q(s, a) ← Q(s, a) + α * [r + γ * max(Q(s', a')) - Q(s, a)]

其中:
α(学习率):控制新信息覆盖旧记忆的速度,通常取0.1~0.3
γ(折扣因子):衡量未来奖励的重要程度,通常取0.8~0.99
r:执行动作后获得的即时奖励
max(Q(s’, a’)):下一状态所有动作中的最大Q值

ε-贪心探索策略

纯粹选择当前Q值最高的动作(贪心策略)容易陷入局部最优。ε-贪心策略以概率ε随机探索新动作,以概率(1-ε)选择当前最优动作,在探索与利用之间取得平衡。训练初期ε较大以充分探索,后期逐渐衰减以稳定策略。

迷宫环境建模

我们将迷宫建模为一个二维网格,其中:
0 表示可通行路径
1 表示墙壁(障碍物)
S 表示起点(Start)
G 表示终点/出口(Goal)

智能体的状态由其所在的网格坐标 (row, col) 唯一确定。每个状态下最多有4个可选动作:向上、向下、向左、向右。若动作会导致撞墙或越界,则智能体位置不变,并给予负奖励作为惩罚。

Java完整实现

项目结构

src/
├── MazeEnvironment.java    // 迷宫环境:状态转移与奖励计算
├── QLearningAgent.java     // Q-Learning智能体:Q-Table与策略
├── MazeTrainer.java        // 训练器:控制训练流程与参数
└── MazeVisualizer.java     // 可视化:打印训练结果与策略路径

1. 迷宫环境类

import java.util.Random;

/**
 * 迷宫环境类
 * 负责定义迷宫地图、执行状态转移、计算即时奖励
 */
public class MazeEnvironment {
    // 地图元素定义
    public static final int PATH = 0;   // 可通行
    public static final int WALL = 1;   // 墙壁
    public static final int START = 2;  // 起点
    public static final int GOAL = 3;   // 终点

    // 动作定义:上、右、下、左(顺时针)
    public static final int ACTION_UP = 0;
    public static final int ACTION_RIGHT = 1;
    public static final int ACTION_DOWN = 2;
    public static final int ACTION_LEFT = 3;
    public static final int ACTION_COUNT = 4;

    // 方向偏移量:{行变化, 列变化}
    private static final int[][] DIRECTIONS = {
        {-1, 0},  // 上
        {0, 1},   // 右
        {1, 0},   // 下
        {0, -1}   // 左
    };

    // 奖励值设计
    private static final double REWARD_GOAL = 100.0;      // 到达终点的奖励
    private static final double REWARD_STEP = -1.0;       // 每步的惩罚(鼓励最短路径)
    private static final double REWARD_HIT_WALL = -5.0;   // 撞墙惩罚

    private final int[][] maze;     // 迷宫地图
    private final int rows;
    private final int cols;
    private final int startRow;
    private final int startCol;
    private final int goalRow;
    private final int goalCol;

    // 当前状态
    private int currentRow;
    private int currentCol;
    private boolean reachedGoal;
    private int steps;

    public MazeEnvironment(int[][] maze) {
        this.maze = maze;
        this.rows = maze.length;
        this.cols = maze[0].length;

        // 查找起点和终点位置
        int[] start = findPosition(START);
        int[] goal = findPosition(GOAL);
        this.startRow = start[0];
        this.startCol = start[1];
        this.goalRow = goal[0];
        this.goalCol = goal[1];

        reset();
    }

    /**
     * 在迷宫中查找指定元素的位置
     */
    private int[] findPosition(int target) {
        for (int r = 0; r < rows; r++) {
            for (int c = 0; c < cols; c++) {
                if (maze[r][c] == target) {
                    return new int[]{r, c};
                }
            }
        }
        throw new IllegalArgumentException("迷宫中未找到目标元素: " + target);
    }

    /**
     * 重置环境到初始状态
     */
    public void reset() {
        this.currentRow = startRow;
        this.currentCol = startCol;
        this.reachedGoal = false;
        this.steps = 0;
    }

    /**
     * 执行动作,返回 {新状态ID, 即时奖励, 是否终止}
     */
    public StepResult step(int action) {
        if (reachedGoal) {
            throw new IllegalStateException("已到达终点,请先调用reset()");
        }

        int newRow = currentRow + DIRECTIONS[action][0];
        int newCol = currentCol + DIRECTIONS[action][1];

        double reward;
        boolean isWall = false;

        // 检查边界和墙壁
        if (newRow < 0 || newRow >= rows || newCol < 0 || newCol >= cols
                || maze[newRow][newCol] == WALL) {
            // 撞墙或越界:位置不变,给予惩罚
            reward = REWARD_HIT_WALL;
            isWall = true;
        } else {
            // 合法移动
            currentRow = newRow;
            currentCol = newCol;
            steps++;

            if (currentRow == goalRow && currentCol == goalCol) {
                reward = REWARD_GOAL;
                reachedGoal = true;
            } else {
                reward = REWARD_STEP;
            }
        }

        int stateId = getStateId(currentRow, currentCol);
        return new StepResult(stateId, reward, reachedGoal, isWall);
    }

    /**
     * 将二维坐标转换为一维状态ID,用于Q-Table索引
     */
    public int getStateId(int row, int col) {
        return row * cols + col;
    }

    /**
     * 将一维状态ID还原为二维坐标
     */
    public int[] getPositionFromStateId(int stateId) {
        return new int[]{stateId / cols, stateId % cols};
    }

    /**
     * 获取状态总数
     */
    public int getStateCount() {
        return rows * cols;
    }

    public int getActionCount() {
        return ACTION_COUNT;
    }

    public int getCurrentStateId() {
        return getStateId(currentRow, currentCol);
    }

    public boolean isReachedGoal() {
        return reachedGoal;
    }

    public int getSteps() {
        return steps;
    }

    public int getGoalRow() { return goalRow; }
    public int getGoalCol() { return goalCol; }
    public int getRows() { return rows; }
    public int getCols() { return cols; }
    public int[][] getMaze() { return maze; }

    /**
     * 单步结果封装类
     */
    public static class StepResult {
        public final int stateId;
        public final double reward;
        public final boolean done;
        public final boolean hitWall;

        public StepResult(int stateId, double reward, boolean done, boolean hitWall) {
            this.stateId = stateId;
            this.reward = reward;
            this.done = done;
            this.hitWall = hitWall;
        }
    }
}

2. Q-Learning智能体类

import java.util.Random;

/**
 * Q-Learning智能体
 * 维护Q-Table,实现ε-贪心策略与Q值更新
 */
public class QLearningAgent {
    private final double[][] qTable;      // Q-Table: [state][action]
    private final int stateCount;
    private final int actionCount;

    // 超参数
    private double alpha;     // 学习率
    private double gamma;     // 折扣因子
    private double epsilon;   // 探索率
    private final double epsilonDecay;    // ε衰减率
    private final double epsilonMin;      // ε最小值

    private final Random random;

    public QLearningAgent(int stateCount, int actionCount,
                          double alpha, double gamma,
                          double epsilon, double epsilonDecay, double epsilonMin) {
        this.stateCount = stateCount;
        this.actionCount = actionCount;
        this.alpha = alpha;
        this.gamma = gamma;
        this.epsilon = epsilon;
        this.epsilonDecay = epsilonDecay;
        this.epsilonMin = epsilonMin;
        this.random = new Random(42);  // 固定随机种子,确保结果可复现

        // 初始化Q-Table为0
        this.qTable = new double[stateCount][actionCount];
    }

    /**
     * ε-贪心策略:选择动作
     * 以ε概率随机探索,以(1-ε)概率选择当前最优动作
     */
    public int selectAction(int stateId) {
        if (random.nextDouble() < epsilon) {
            // 探索:随机选择动作
            return random.nextInt(actionCount);
        } else {
            // 利用:选择Q值最高的动作
            return getBestAction(stateId);
        }
    }

    /**
     * 获取指定状态下Q值最高的动作
     */
    public int getBestAction(int stateId) {
        int bestAction = 0;
        double maxQ = qTable[stateId][0];

        for (int a = 1; a < actionCount; a++) {
            if (qTable[stateId][a] > maxQ) {
                maxQ = qTable[stateId][a];
                bestAction = a;
            }
        }
        return bestAction;
    }

    /**
     * 获取指定状态下所有动作中的最大Q值
     */
    public double getMaxQ(int stateId) {
        double maxQ = qTable[stateId][0];
        for (int a = 1; a < actionCount; a++) {
            if (qTable[stateId][a] > maxQ) {
                maxQ = qTable[stateId][a];
            }
        }
        return maxQ;
    }

    /**
     * 根据贝尔曼方程更新Q值
     */
    public void update(int stateId, int action, double reward, int nextStateId, boolean done) {
        double currentQ = qTable[stateId][action];
        double targetQ;

        if (done) {
            // 终止状态:没有未来奖励
            targetQ = reward;
        } else {
            // 非终止状态:加上折扣后的未来最大Q值
            targetQ = reward + gamma * getMaxQ(nextStateId);
        }

        // 向目标Q值逐步逼近(时序差分学习)
        qTable[stateId][action] = currentQ + alpha * (targetQ - currentQ);
    }

    /**
     * 衰减探索率,训练后期减少随机探索
     */
    public void decayEpsilon() {
        if (epsilon > epsilonMin) {
            epsilon *= epsilonDecay;
            if (epsilon < epsilonMin) {
                epsilon = epsilonMin;
            }
        }
    }

    /**
     * 获取当前Q-Table(用于可视化)
     */
    public double[][] getQTable() {
        return qTable;
    }

    public double getEpsilon() {
        return epsilon;
    }

    public double getAlpha() {
        return alpha;
    }

    public double getGamma() {
        return gamma;
    }
}

3. 训练器类

/**
 * 训练控制器
 * 协调环境与智能体的交互,执行多轮训练episode
 */
public class MazeTrainer {
    private final MazeEnvironment env;
    private final QLearningAgent agent;
    private final int maxStepsPerEpisode;  // 每轮最大步数,防止无限循环

    // 训练统计
    private int[] episodeSteps;
    private double[] episodeRewards;
    private int[] episodeSuccess;

    public MazeTrainer(MazeEnvironment env, QLearningAgent agent, int maxStepsPerEpisode) {
        this.env = env;
        this.agent = agent;
        this.maxStepsPerEpisode = maxStepsPerEpisode;
    }

    /**
     * 执行完整训练
     * @param episodes 训练轮数
     */
    public void train(int episodes) {
        episodeSteps = new int[episodes];
        episodeRewards = new double[episodes];
        episodeSuccess = new int[episodes];

        System.out.println("开始训练,总轮数: " + episodes);
        System.out.println("初始参数: α=" + agent.getAlpha()
                + ", γ=" + agent.getGamma()
                + ", ε=" + agent.getEpsilon());
        System.out.println("-".repeat(60));

        for (int ep = 0; ep < episodes; ep++) {
            env.reset();
            int stateId = env.getCurrentStateId();
            double totalReward = 0;
            int steps = 0;
            boolean success = false;

            while (steps < maxStepsPerEpisode) {
                // 1. 根据当前策略选择动作
                int action = agent.selectAction(stateId);

                // 2. 执行动作,观察环境和奖励
                MazeEnvironment.StepResult result = env.step(action);

                // 3. 更新Q值
                agent.update(stateId, action, result.reward, result.stateId, result.done);

                totalReward += result.reward;
                steps++;
                stateId = result.stateId;

                if (result.done) {
                    success = true;
                    break;
                }
            }

            // 4. 每轮结束后衰减探索率
            agent.decayEpsilon();

            episodeSteps[ep] = steps;
            episodeRewards[ep] = totalReward;
            episodeSuccess[ep] = success ? 1 : 0;

            // 每100轮打印一次进度
            if ((ep + 1) % 100 == 0) {
                double recentSuccess = calculateSuccessRate(ep, 100);
                double recentReward = calculateAvgReward(ep, 100);
                System.out.printf("轮次 %d | 近100轮成功率: %.1f%% | 平均奖励: %.2f | ε: %.4f%n",
                        ep + 1, recentSuccess * 100, recentReward, agent.getEpsilon());
            }
        }

        System.out.println("-".repeat(60));
        System.out.println("训练完成!");
    }

    /**
     * 计算最近N轮的成功率
     */
    private double calculateSuccessRate(int currentEp, int window) {
        int start = Math.max(0, currentEp - window + 1);
        int sum = 0;
        for (int i = start; i <= currentEp; i++) {
            sum += episodeSuccess[i];
        }
        return (double) sum / (currentEp - start + 1);
    }

    /**
     * 计算最近N轮的平均奖励
     */
    private double calculateAvgReward(int currentEp, int window) {
        int start = Math.max(0, currentEp - window + 1);
        double sum = 0;
        for (int i = start; i <= currentEp; i++) {
            sum += episodeRewards[i];
        }
        return sum / (currentEp - start + 1);
    }

    public int[] getEpisodeSteps() { return episodeSteps; }
    public double[] getEpisodeRewards() { return episodeRewards; }
    public int[] getEpisodeSuccess() { return episodeSuccess; }
}

4. 可视化与主程序

import java.util.ArrayList;
import java.util.List;

/**
 * 可视化工具与主入口
 * 打印迷宫、策略路径和训练统计
 */
public class MazeVisualizer {

    /**
     * 打印迷宫地图,标记智能体当前位置
     */
    public static void printMaze(MazeEnvironment env, int agentRow, int agentCol) {
        int[][] maze = env.getMaze();
        for (int r = 0; r < env.getRows(); r++) {
            for (int c = 0; c < env.getCols(); c++) {
                if (r == agentRow && c == agentCol) {
                    System.out.print(" A ");  // Agent位置
                } else if (maze[r][c] == MazeEnvironment.WALL) {
                    System.out.print("███");
                } else if (maze[r][c] == MazeEnvironment.GOAL) {
                    System.out.print(" G ");
                } else if (maze[r][c] == MazeEnvironment.START) {
                    System.out.print(" S ");
                } else {
                    System.out.print("   ");
                }
            }
            System.out.println();
        }
    }

    /**
     * 根据训练好的Q-Table,提取从起点到终点的最优路径
     */
    public static List<int[]> extractOptimalPath(MazeEnvironment env, QLearningAgent agent) {
        List<int[]> path = new ArrayList<>();
        env.reset();
        int stateId = env.getCurrentStateId();
        int[] pos = env.getPositionFromStateId(stateId);
        path.add(new int[]{pos[0], pos[1]});

        int maxSteps = env.getRows() * env.getCols();  // 防止循环
        int steps = 0;

        while (steps < maxSteps) {
            int action = agent.getBestAction(stateId);
            MazeEnvironment.StepResult result = env.step(action);
            stateId = result.stateId;
            pos = env.getPositionFromStateId(stateId);
            path.add(new int[]{pos[0], pos[1]});

            if (result.done) {
                break;
            }
            steps++;
        }

        return path;
    }

    /**
     * 打印最优路径在迷宫上的可视化
     */
    public static void printPathOnMaze(MazeEnvironment env, List<int[]> path) {
        int[][] maze = env.getMaze();
        boolean[][] onPath = new boolean[env.getRows()][env.getCols()];
        for (int[] p : path) {
            onPath[p[0]][p[1]] = true;
        }

        System.out.println("\n最优路径(*表示路径,A→终点):");
        for (int r = 0; r < env.getRows(); r++) {
            for (int c = 0; c < env.getCols(); c++) {
                if (maze[r][c] == MazeEnvironment.WALL) {
                    System.out.print("███");
                } else if (maze[r][c] == MazeEnvironment.GOAL) {
                    System.out.print(" G ");
                } else if (maze[r][c] == MazeEnvironment.START) {
                    System.out.print(" S ");
                } else if (onPath[r][c]) {
                    System.out.print(" * ");
                } else {
                    System.out.print("   ");
                }
            }
            System.out.println();
        }
        System.out.println("路径长度: " + (path.size() - 1) + " 步");
    }

    public static void main(String[] args) {
        // ========== 定义迷宫地图 ==========
        // 0=通路, 1=墙壁, 2=起点, 3=终点
        int[][] mazeMap = {
            {1, 1, 1, 1, 1, 1, 1, 1, 1, 1},
            {1, 2, 0, 0, 1, 0, 0, 0, 0, 1},
            {1, 1, 1, 0, 1, 0, 1, 1, 0, 1},
            {1, 0, 0, 0, 0, 0, 1, 0, 0, 1},
            {1, 0, 1, 1, 1, 1, 1, 0, 1, 1},
            {1, 0, 0, 0, 0, 0, 0, 0, 0, 1},
            {1, 1, 1, 0, 1, 1, 1, 1, 0, 1},
            {1, 0, 1, 0, 0, 0, 1, 0, 0, 1},
            {1, 0, 0, 0, 1, 0, 0, 0, 3, 1},
            {1, 1, 1, 1, 1, 1, 1, 1, 1, 1}
        };

        // ========== 初始化环境与智能体 ==========
        MazeEnvironment env = new MazeEnvironment(mazeMap);
        QLearningAgent agent = new QLearningAgent(
            env.getStateCount(),      // 状态数
            env.getActionCount(),     // 动作数
            0.2,                      // α: 学习率
            0.95,                     // γ: 折扣因子(高度重视未来奖励)
            1.0,                      // ε: 初始探索率(100%随机探索)
            0.995,                    // ε衰减率
            0.01                      // ε最小值
        );

        // ========== 执行训练 ==========
        MazeTrainer trainer = new MazeTrainer(env, agent, maxStepsPerEpisode: 200);
        trainer.train(2000);  // 训练2000轮

        // ========== 展示训练结果 ==========
        System.out.println("\n" + "=".repeat(60));
        System.out.println("训练完成后,使用贪心策略提取最优路径:");
        System.out.println("=".repeat(60));

        List<int[]> optimalPath = extractOptimalPath(env, agent);
        printPathOnMaze(env, optimalPath);

        // 打印最终成功率
        int totalSuccess = 0;
        for (int s : trainer.getEpisodeSuccess()) {
            totalSuccess += s;
        }
        System.out.printf("\n总体成功率: %.1f%%%n", (double) totalSuccess / 2000 * 100);
    }
}

训练过程分析与参数调优

超参数影响

参数 作用 推荐值 调优建议
学习率 α 控制Q值更新步长 0.1~0.3 过大导致震荡,过小收敛慢
折扣因子 γ 未来奖励权重 0.9~0.99 迷宫问题建议0.95以上,重视长远规划
初始探索率 ε 随机动作比例 1.0 初期充分探索环境
ε衰减率 探索→利用的过渡速度 0.995~0.999 总轮数多时可调小衰减
ε最小值 保留的最低探索比例 0.01~0.05 防止陷入局部最优

收敛特征

在10×10的迷宫上训练2000轮,典型收敛曲线如下:
前200轮:成功率低于20%,智能体以随机探索为主,频繁撞墙或超时
200~800轮:成功率快速上升至70%~90%,Q-Table开始形成有效策略
800轮以后:成功率稳定在95%以上,ε已衰减至较低水平,策略基本收敛

奖励塑形技巧

上述实现中,每步给予-1的微小惩罚,到达终点给予+100的大奖励。这种设计隐式地引导智能体寻找最短路径——路径越长,累积的步数惩罚越大。若迷宫存在多条等长路径,智能体可能收敛至其中任意一条。

算法复杂度分析

  • 时间复杂度:每轮训练为 O(maxSteps),总训练时间为 O(episodes × maxSteps)。Q-Table查询和更新均为 O(1)(数组随机访问)。
  • 空间复杂度O(stateCount × actionCount),用于存储Q-Table。对于N×M的迷宫,空间复杂度为 O(N×M×4),即 O(N×M),非常紧凑。
  • 与A*等传统路径规划对比:Q-Learning不需要预先知道地图结构(免模型),在动态环境(如障碍物随机移动)中优势显著;但在静态已知环境中,A*的计算效率更高。

扩展方向

  1. 深度Q网络(DQN):当状态空间极大时(如高分辨率图像输入),用神经网络替代Q-Table,实现端到端的感知-决策
  2. SARSA算法:与Q-Learning的离策略(Off-Policy)不同,SARSA采用同策略(On-Policy)更新,对探索过程更保守,在某些环境中更安全
  3. 多智能体协作:多个智能体共享Q-Table或独立学习,在复杂迷宫中分工探索
  4. 连续动作空间:结合策略梯度方法(如REINFORCE),将Q-Learning扩展到连续控制领域

总结

本文从马尔可夫决策过程的形式化定义出发,完整地用Java实现了Q-Learning算法在迷宫探索中的应用。关键要点包括:

  • 状态空间离散化:将二维迷宫坐标映射为一维状态ID,便于Q-Table索引
  • ε-贪心策略:在训练早期鼓励探索,后期稳定至贪心策略,兼顾学习效率与最终性能
  • 奖励塑形:通过步数惩罚和终点奖励的精心设计,隐式引导最短路径行为
  • 完整工程实现:环境、智能体、训练器、可视化四层架构清晰分离,易于扩展

强化学习的魅力在于”授人以渔”——我们无需告诉智能体具体的路径,只需定义好奖励规则,它就能通过试错自主发现最优策略。这种自主学习的能力,正是AI从”工具”走向”智能”的关键一步。