本项目是对西湖大学赵世钰教授强化学习课程的完整复现与实现,涵盖了从基础的动态规划到深度强化学习的各类经典算法。所有算法均在网格世界(GridWorld)环境中实现和验证,提供了清晰的可视化结果和详细的中文注释。
本项目实现了强化学习领域的 15+ 种经典算法,包括:
- 基于模型的方法:值迭代、策略迭代、截断策略迭代
- 蒙特卡洛方法:MC Exploring Starts
- 时序差分学习:Sarsa、Expected Sarsa、n-step Sarsa、Q-Learning
- 值函数逼近:线性 Sarsa、线性 Q-Learning、Deep Q-Network (DQN)
- 策略梯度方法:REINFORCE、QAC、A2C、Off-Policy AC、DPG
所有算法都经过精心设计和优化,具有:
- ✅ 完整的中文注释
- ✅ 统一的代码架构
- ✅ 清晰的可视化结果
- ✅ 详细的训练过程记录
项目提供了完整的 Web 前端界面,支持在线交互训练:
- 环境配置面板:可视化配置网格大小、障碍物布局,支持预设迷宫模板(开关式,每次随机生成)
- 交互式障碍物编辑:直接点击网格放置/移除障碍物,起点和终点不可放置
- 路径合法性验证:BFS 自动检测是否存在从起点到终点的可行路径,防止无效配置
- 6 种算法可选:值迭代、策略迭代、截断策略迭代、Sarsa、Q-Learning、MC Exploring Starts
- 实时训练看板:WebSocket 实时推送训练进度、策略热力图、回报曲线和收敛曲线
- 可折叠面板:环境配置、算法选择、超参数、高级选项均支持折叠
启动方式:
终端 1: cd backend && python main.py # 后端 :8000
终端 2: cd frontend && npm run dev # 前端 :5173
项目采用高度模块化的架构,将通用功能抽象为独立模块:
- 特征工程模块:统一的状态特征提取
- 策略选择模块:贪心策略、ε-greedy 策略等
- 神经网络模块:Actor、Critic、Q-Network、DQN 等
- 评估指标模块:滑动平均、策略变化统计等
- 环境配置模块:预定义的网格环境配置
所有算法的训练结果统一保存到 outputs/ 目录,按算法类别组织:
outputs/
├── Value-iteration-and-Policy-iteration/
├── Monte-Carlo-Learning/
├── Temporal-Difference-Learning/
├── Value-Function-Approximation/
└── Policy-Gradient-Methods/
与传统实现不同,本项目中的障碍物默认可穿越但会给予负奖励惩罚;同时支持不可穿越模式(选中后碰到障碍物原地不动)。这种设计:
- 更符合真实场景(如泥泞地形、收费路段)
- 增加了策略学习的复杂度
- 提供了更丰富的探索空间
每个算法都提供 2×2 的可视化结果图:
- 策略热力图:显示最终策略和状态值函数
- 训练回报曲线:展示学习过程和收敛情况
- 策略变化曲线:追踪策略的演化过程
- 动作价值分布:分析起点状态的动作选择
- Python 3.8+
- PyTorch 2.0+
- NumPy
- Matplotlib
- tqdm
- Node.js 18+(Web 前端)
# 克隆仓库
git clone https://github.com/21661/reinforcement-learning-course.git
cd reinforcement-learning-course
# Python 依赖
pip install torch numpy matplotlib tqdm
# 后端依赖
cd backend && pip install -r requirements.txt && cd ..
# 前端依赖
cd frontend && npm install && cd ..Web 前端(推荐):
# 终端 1:启动后端
cd backend && python main.py
# 终端 2:启动前端
cd frontend && npm run dev打开浏览器访问 http://localhost:5173 ,在左侧面板选择算法、配置环境参数,点击"保存配置"后即可"开始训练"。右侧实时显示策略热力图和训练回报曲线。
命令行运行:
# 运行 DPG 算法
python src/Policy-Gradient-Methods/Action-Cirtic-Methods/dpg.py
# 运行 Sarsa 算法
python src/Temporal-Difference\ Learing/sarsa/sarsa_basic.py
# 运行 DQN 算法
python src/Value-Function-Approximation/deep-Q-learning/deep_q_learning.py| 算法 | 文件路径 | 说明 |
|---|---|---|
| 值迭代 | src/Value-iteration-and-Policy-iteration/value_iteration.py |
通过迭代更新状态值函数求解最优策略 |
| 策略迭代 | src/Value-iteration-and-Policy-iteration/policy_iteration.py |
策略评估与策略改进交替进行 |
| 截断策略迭代 | src/Value-iteration-and-Policy-iteration/truncated_policy_iteration.py |
限制策略评估的迭代次数 |
| 算法 | 文件路径 | 说明 |
|---|---|---|
| MC Exploring Starts | src/Monte-Carlo-Learing/mc_exploring_starts.py |
通过完整回合采样学习动作值函数 |
| 算法 | 文件路径 | 说明 |
|---|---|---|
| Sarsa | src/Temporal-Difference Learing/sarsa/sarsa_basic.py |
On-policy TD 控制算法 |
| Expected Sarsa | src/Temporal-Difference Learing/sarsa/sarsa_basic.py |
使用期望更新的 Sarsa 变体 |
| n-step Sarsa | src/Temporal-Difference Learing/sarsa/sarsa_basic.py |
多步自举的 Sarsa 算法 |
| Q-Learning | src/Temporal-Difference Learing/Q-learning/q_learning.py |
Off-policy TD 控制算法 |
| 算法 | 文件路径 | 说明 |
|---|---|---|
| 线性 Sarsa | src/Value-Function-Approximation/sarsa/sarsa_linear.py |
使用线性函数逼近的 Sarsa |
| 线性 Q-Learning | src/Value-Function-Approximation/Q-learning/q_learning_linear.py |
使用线性函数逼近的 Q-Learning |
| DQN | src/Value-Function-Approximation/deep-Q-learning/deep_q_learning.py |
深度 Q 网络,使用经验回放和目标网络 |
| 算法 | 文件路径 | 说明 |
|---|---|---|
| REINFORCE | src/Policy-Gradient-Methods/RAINFORCE/reinforce.py |
蒙特卡洛策略梯度算法 |
| QAC | src/Policy-Gradient-Methods/Action-Cirtic-Methods/qac.py |
Q Actor-Critic,使用 Q 函数作为 Critic |
| A2C | src/Policy-Gradient-Methods/Action-Cirtic-Methods/a2c.py |
Advantage Actor-Critic,使用优势函数 |
| Off-Policy AC | src/Policy-Gradient-Methods/Action-Cirtic-Methods/off_policy_actor_critic.py |
离策略 Actor-Critic |
| DPG | src/Policy-Gradient-Methods/Action-Cirtic-Methods/dpg.py |
确定性策略梯度算法 |
RL-Algorithms-Implementation/
├── backend/ # Web 后端(FastAPI + WebSocket)
│ ├── main.py # 训练 API + WebSocket 端点
│ ├── trainer.py # 6 种算法生成器适配
│ └── requirements.txt
│
├── frontend/ # Web 前端(React + Vite + TypeScript)
│ └── src/
│ ├── App.tsx # 主页面(三栏布局)
│ ├── types.ts # TypeScript 类型
│ ├── hooks/useWebSocket.ts # WebSocket 连接管理
│ └── components/
│ ├── ConfigPanel.tsx # 环境配置 + 算法选择 + 超参数
│ ├── GridWorldView.tsx # 网格热力图 + 可点击编辑障碍物
│ └── TrainingDashboard.tsx # 回报曲线 + 收敛曲线 + 进度条
│
├── src/
│ ├── rl_utils/ # 统一工具模块
│ │ ├── grid_world.py # 网格世界环境(支持不可穿越模式)
│ │ ├── visualize.py # 可视化工具
│ │ ├── features.py # 特征工程
│ │ ├── policies.py # 策略选择
│ │ ├── networks.py # 神经网络
│ │ ├── metrics.py # 评估指标
│ │ └── configs.py # 环境配置
│ │
│ ├── Value-iteration-and-Policy-iteration/
│ ├── Monte-Carlo-Learing/
│ ├── Temporal-Difference Learing/
│ ├── Value-Function-Approximation/
│ └── Policy-Gradient-Methods/
│
├── outputs/ # 训练结果输出
├── scripts/ # 工具脚本
└── README.md
from rl_utils import GridWorld, OBSTACLES_6X6
# 使用预定义配置
env = GridWorld(
rows=6, cols=6,
start=(0, 0), goal=(5, 5),
obstacles=OBSTACLES_6X6,
gamma=0.9,
reward_step=-1,
reward_goal=100,
reward_obstacle=-10,
slip_prob=0.2
)from rl_utils import (
state_to_feature,
ActorNetwork,
QNetwork,
moving_average,
choose_greedy_action
)
# 特征提取
features = state_to_feature(env, state)
# 创建神经网络
actor = ActorNetwork(feature_dim=20, n_actions=4)
q_net = QNetwork(feature_dim=20, n_actions=4)
# 策略选择
action = choose_greedy_action(env, Q, state)
# 计算滑动平均
smoothed = moving_average(returns, window_size=200)所有算法都在 6×6 或 10×10 的网格世界中进行了充分训练和测试。训练结果包括:
- 收敛性验证:所有算法都能成功收敛到最优或近似最优策略
- 学习效率对比:不同算法的样本效率和收敛速度对比
- 策略质量分析:最终策略的回报和路径长度统计
详细的实验结果和可视化图表保存在 outputs/ 目录中。
- 可配置的网格大小:支持任意大小的网格
- 灵活的障碍物设置:默认可穿越但有惩罚,支持切换为不可穿越模式(
obstacle_impassable=True) - 随机性支持:可配置的滑动概率(slip_prob)
- 自定义奖励函数:支持设置不同的奖励值
- FastAPI 后端:REST API + WebSocket 实时推送训练进度
- React 前端:TypeScript + Vite + Tailwind CSS + Recharts
- 交互式障碍物编辑:点击网格直接放置/移除障碍物
- 路径验证:BFS 检查起点到终点的可达性
- 实时可视化:策略热力图、回报曲线、收敛曲线
- 策略热力图:显示状态值函数和策略箭头
- 训练曲线:回报曲线、策略变化曲线
- 动作价值分布:柱状图展示 Q 值分布
- 障碍物标识:半透明红色标识可穿越障碍物
- ActorNetwork:随机策略网络(用于 REINFORCE、QAC、A2C)
- DeterministicActor:确定性策略网络(用于 DPG)
- ValueNetwork:状态值函数网络(用于 A2C)
- QNetwork:动作值函数网络(用于 QAC、DPG)
- DQN:深度 Q 网络(用于 DQN)
-
Sutton & Barto - Reinforcement Learning: An Introduction (2nd Edition)
- 强化学习领域的经典教材
- 本项目的理论基础
-
西湖大学赵世钰教授课程
- 本项目的课程来源
- 系统的强化学习理论讲解
欢迎提交 Issue 和 Pull Request!
- Fork 本仓库
- 创建特性分支 (
git checkout -b feature/AmazingFeature) - 提交更改 (
git commit -m 'Add some AmazingFeature') - 推送到分支 (
git push origin feature/AmazingFeature) - 开启 Pull Request
- 遵循 PEP 8 代码风格
- 提供完整的中文注释
- 添加必要的单元测试
- 更新相关文档
- ✅ 新增 Web 在线训练平台(FastAPI 后端 + React 前端)
- ✅ WebSocket 实时推送训练进度、策略热力图、回报曲线
- ✅ 交互式障碍物编辑(点击网格放置/移除)
- ✅ 路径合法性 BFS 验证
- ✅ 环境支持不可穿越障碍物模式(
obstacle_impassable) - ✅ 可折叠配置面板(算法选择、超参数、高级选项)
- ✅ 实现 15+ 种强化学习算法
- ✅ 完成项目重构,消除重复代码
- ✅ 统一输出管理和可视化
- ✅ 添加完整的中文文档
- 感谢西湖大学赵世钰教授的精彩课程
- 感谢 Sutton & Barto 的经典教材
- 感谢开源社区的支持和贡献
⭐ 如果这个项目对你有帮助,欢迎 Star 支持!
