Skip to content

Norlaner/reinforcement-learning-course

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

6 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

强化学习算法实现与复现

Python 3.8+ PyTorch

本项目是对西湖大学赵世钰教授强化学习课程的完整复现与实现,涵盖了从基础的动态规划到深度强化学习的各类经典算法。所有算法均在网格世界(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

所有算法都经过精心设计和优化,具有:

  • ✅ 完整的中文注释
  • ✅ 统一的代码架构
  • ✅ 清晰的可视化结果
  • ✅ 详细的训练过程记录

🎯 项目特色

1. Web 在线训练平台(NEW)

项目提供了完整的 Web 前端界面,支持在线交互训练:

  • 环境配置面板:可视化配置网格大小、障碍物布局,支持预设迷宫模板(开关式,每次随机生成)
  • 交互式障碍物编辑:直接点击网格放置/移除障碍物,起点和终点不可放置
  • 路径合法性验证:BFS 自动检测是否存在从起点到终点的可行路径,防止无效配置
  • 6 种算法可选:值迭代、策略迭代、截断策略迭代、Sarsa、Q-Learning、MC Exploring Starts
  • 实时训练看板:WebSocket 实时推送训练进度、策略热力图、回报曲线和收敛曲线
  • 可折叠面板:环境配置、算法选择、超参数、高级选项均支持折叠

界面展示

启动方式:
  终端 1: cd backend && python main.py     # 后端 :8000
  终端 2: cd frontend && npm run dev       # 前端 :5173

2. 模块化设计

项目采用高度模块化的架构,将通用功能抽象为独立模块:

  • 特征工程模块:统一的状态特征提取
  • 策略选择模块:贪心策略、ε-greedy 策略等
  • 神经网络模块:Actor、Critic、Q-Network、DQN 等
  • 评估指标模块:滑动平均、策略变化统计等
  • 环境配置模块:预定义的网格环境配置

3. 统一的输出管理

所有算法的训练结果统一保存到 outputs/ 目录,按算法类别组织:

outputs/
├── Value-iteration-and-Policy-iteration/
├── Monte-Carlo-Learning/
├── Temporal-Difference-Learning/
├── Value-Function-Approximation/
└── Policy-Gradient-Methods/

4. 可穿越障碍物设计

与传统实现不同,本项目中的障碍物默认可穿越但会给予负奖励惩罚;同时支持不可穿越模式(选中后碰到障碍物原地不动)。这种设计:

  • 更符合真实场景(如泥泞地形、收费路段)
  • 增加了策略学习的复杂度
  • 提供了更丰富的探索空间

5. 丰富的可视化

每个算法都提供 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

📖 算法列表

1. 基于模型的方法(Model-Based Methods)

算法 文件路径 说明
值迭代 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 限制策略评估的迭代次数

2. 蒙特卡洛方法(Monte Carlo Methods)

算法 文件路径 说明
MC Exploring Starts src/Monte-Carlo-Learing/mc_exploring_starts.py 通过完整回合采样学习动作值函数

3. 时序差分学习(Temporal Difference Learning)

算法 文件路径 说明
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 控制算法

4. 值函数逼近(Value Function Approximation)

算法 文件路径 说明
线性 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 网络,使用经验回放和目标网络

5. 策略梯度方法(Policy Gradient Methods)

算法 文件路径 说明
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)
  • 自定义奖励函数:支持设置不同的奖励值

Web 在线训练平台

  • 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)

🎓 学习资源

推荐阅读

  1. Sutton & Barto - Reinforcement Learning: An Introduction (2nd Edition)

    • 强化学习领域的经典教材
    • 本项目的理论基础
  2. 西湖大学赵世钰教授课程

    • 本项目的课程来源
    • 系统的强化学习理论讲解

相关链接

🤝 贡献指南

欢迎提交 Issue 和 Pull Request!

贡献方式

  1. Fork 本仓库
  2. 创建特性分支 (git checkout -b feature/AmazingFeature)
  3. 提交更改 (git commit -m 'Add some AmazingFeature')
  4. 推送到分支 (git push origin feature/AmazingFeature)
  5. 开启 Pull Request

代码规范

  • 遵循 PEP 8 代码风格
  • 提供完整的中文注释
  • 添加必要的单元测试
  • 更新相关文档

📝 更新日志

v1.1.0 (2026-07-20)

  • ✅ 新增 Web 在线训练平台(FastAPI 后端 + React 前端)
  • ✅ WebSocket 实时推送训练进度、策略热力图、回报曲线
  • ✅ 交互式障碍物编辑(点击网格放置/移除)
  • ✅ 路径合法性 BFS 验证
  • ✅ 环境支持不可穿越障碍物模式(obstacle_impassable
  • ✅ 可折叠配置面板(算法选择、超参数、高级选项)

v1.0.0 (2026-03-31)

  • ✅ 实现 15+ 种强化学习算法
  • ✅ 完成项目重构,消除重复代码
  • ✅ 统一输出管理和可视化
  • ✅ 添加完整的中文文档

🙏 致谢

  • 感谢西湖大学赵世钰教授的精彩课程
  • 感谢 Sutton & Barto 的经典教材
  • 感谢开源社区的支持和贡献

⭐ 如果这个项目对你有帮助,欢迎 Star 支持!

About

西湖大学赵世钰教授强化学习课程的完整复现与实现,涵盖15+种经典算法,包括值迭代、蒙特卡洛、时序差分、值函数逼近和策略梯度方法。

Resources

License

Stars

11 stars

Watchers

0 watching

Forks

Releases

No releases published

Packages

 
 
 

Contributors