Skip to content

dsyt-SichenGuo/gsc-reinforcement-learning-course

 
 

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

4 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. 模块化设计

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

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

2. 统一的输出管理

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

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

3. 可穿越障碍物设计

与传统实现不同,本项目中的障碍物是可穿越的,但会给予负奖励惩罚。这种设计:

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

4. 丰富的可视化

每个算法都提供 2×2 的可视化结果图:

  • 策略热力图:显示最终策略和状态值函数
  • 训练回报曲线:展示学习过程和收敛情况
  • 策略变化曲线:追踪策略的演化过程
  • 动作价值分布:分析起点状态的动作选择

🚀 快速开始

环境要求

  • Python 3.8+
  • PyTorch 2.0+
  • NumPy
  • Matplotlib
  • tqdm

安装依赖

# 克隆仓库
git clone https://github.com/21661/reinforcement-learning-course.git
cd reinforcement-learning-course

# 安装依赖
pip install torch numpy matplotlib tqdm

运行示例

# 运行 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/
├── 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/ 目录中。

🔧 核心功能

网格世界环境

  • 可配置的网格大小:支持任意大小的网格
  • 灵活的障碍物设置:可穿越但有惩罚
  • 随机性支持:可配置的滑动概率(slip_prob)
  • 自定义奖励函数:支持设置不同的奖励值

可视化工具

  • 策略热力图:显示状态值函数和策略箭头
  • 训练曲线:回报曲线、策略变化曲线
  • 动作价值分布:柱状图展示 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.0.0 (2026-03-31)

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

🙏 致谢

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

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

About

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

Resources

License

Stars

0 stars

Watchers

0 watching

Forks

Releases

No releases published

Packages

 
 
 

Contributors

Languages

  • Python 100.0%