- 人工智能
- 机器学习
- 深度学习
【免费下载链接】mlpack
mlpack: a fast, header-only C++ machine learning library
mlpack 提供了一套完整的端到端强化学习(Reinforcement Learning, RL)框架,覆盖经典基准环境(Cart Pole、Acrobot、Mountain Car 等)、常用策略、经验回放方法以及 Q-Learning、异步学习、DDPG、TD3、SAC 等多种主流智能体,且支持自定义环境与策略的无运行时开销接入。本文以 doc/tutorials/reinforcement_learning.md 为主线,结合src/mlpack/methods/reinforcement_learning/下的源码实现,系统讲解 mlpack RL 框架的组成、关键参数含义与真实可运行的训练示例,帮助你快速上手用 mlpack 训练自己的强化学习智能体。
mlpack 强化学习框架概览
强化学习是当前机器学习领域最热门的方向之一,自 DeepMind 发表利用深度神经网络训练智能体玩 Atari 游戏并取得巨大成功的论文以来,相关研究与应用持续升温。mlpack 实现了完整的强化学习端到端框架,其核心代码集中在 src/mlpack/methods/reinforcement_learning/,主要包含以下模块:
- 环境(Environment):Cart Pole、Acrobot、Mountain Car、Pendulum 等经典基准测试环境;
- 策略(Policy):Greedy Policy、Aggregated Policy 等行为策略;
- 回放(Replay):Random Replay、Prioritized Replay 等经验回放方法;
- 网络(Q-Networks):SimpleDQN、Dueling DQN、Categorical DQN 等动作价值网络;
- 噪声(Noise):GaussianNoise、OUNoise(Ornstein-Uhlenbeck 过程);
- 智能体(Agent):QLearning、AsyncLearning、DDPG、TD3、SAC。
从源码结构看,reinforcement_learning.hpp 是所有这些组件的统一聚合头文件,包含environment/、policy/、q_networks/、replay/、worker/、noise/六大子模块以及training_config.hpp和四个智能体类。也就是说,只需#include <mlpack.hpp>即可使用全部 RL 能力。
除了内置环境,mlpack 还可以通过 gym_tcp_api 与 OpenAI Gym 工具包通信,从而接入更丰富的环境(如 Atari 游戏)。由于框架基于模板实现,自定义环境与策略同样可以直接接入,且不引入任何运行时开销。
强化学习环境:State、Action 与 Sample
mlpack 实现了多个最流行的强化学习测试环境,包括 Cart Pole、Acrobot、Mountain Car 及其变体。所有这些环境都通过 environment/environment.hpp 统一引入,其中包含:
Acrobot:双连杆摆,仅第二个关节可驱动,目标是摆动末端执行器至指定高度;CartPole:小车倒立摆,状态为(位置、速度、角度、角速度)四元组;ContinuousMountainCar、MountainCar:连续/离散动作的爬山车;DoublePoleCart、ContDoublePoleCart:双杆小车及其连续版本;Pendulum:连续动作空间的钟摆控制(DDPG、TD3、SAC 示例均使用该环境);- 以及
reward_clipping奖励裁剪等辅助工具。
环境可扩展性是 mlpack 框架的关键特性。自定义环境只需实现一组特定的方法与类,即可被智能体无缝使用:
State:环境的表示。以CartPole为例,cart_pole.hpp 中的State类存储位置(Position)、速度(Velocity)、角度(Angle)与角速度(AngularVelocity)四个成员,内部以arma::colvec向量承载,维度为 4,并提供Encode()方法将状态编码为列向量供神经网络使用。Action:对于离散环境,Action是一个带枚举类型的类,枚举列出智能体在环境中可执行的全部动作。继续以CartPole为例,其枚举仅包含backward和forward两个动作,并定义了static const size_t size = 2用于标记动作空间大小;对于连续环境(如Pendulum),Action类则包含一个大小取决于动作空间维度的数组。Sample:这是环境的心脏。CartPole::Sample(const State& state, const Action& action, State& nextState)根据当前状态与动作计算物理动力学,返回奖励并更新下一状态。从源码看,CartPole环境默认返回恒为 1.0 的奖励,直到智能体失败(小车越界或角度超出阈值)或达到最大步数。
自定义环境通常还会用到若干辅助方法,例如Acrobot环境中的Dsdt方法(状态导数),配合RK4迭代方法(四阶龙格-库塔法)来估计下一状态。每个环境还需提供InitialSample()(生成每个回合的初始状态,CartPole中初始状态在 [-0.05, 0.05] 内随机)与IsTerminal()(判断是否到达终止状态)等辅助接口。
CartPole构造函数的完整参数(源码见 cart_pole.hpp)如下:
| 参数 | 默认值 | 含义 |
|---|---|---|
maxSteps | 200 | 回合终止前最大步数,0 表示无限制 |
gravity | 9.8 | 重力常数 |
massCart | 1.0 | 小车质量 |
massPole | 0.1 | 杆的质量 |
length | 0.5 | 杆的长度 |
forceMag | 10.0 | 施加力的大小 |
tau | 0.02 | 时间步间隔 |
thetaThresholdRadians | 12° 对应弧度 | 角度终止阈值 |
xThreshold | 2.4 | 位置终止阈值 |
doneReward | 1.0 | 成功完成回合时获得的奖励 |
强化学习智能体的组成:策略与回放
一个强化学习智能体在环境中采取动作,以最大化累积奖励。为此它需要两样东西:选择动作的方式(策略 Policy)与采样历史经验的方式(回放 Replay)。
策略(Policy)
最简单的策略示例是 epsilon-greedy 策略:智能体以概率 epsilon 随机探索,其余时候贪心选择当前价值最高的动作。epsilon 会随时间逐步衰减,从而在探索(exploration)与利用(exploitation)之间取得平衡。
mlpack 的GreedyPolicy实现位于 policy/greedy_policy.hpp,构造函数签名为:
GreedyPolicy(const double initialEpsilon, const size_t annealInterval, const double minEpsilon, const double decayRate = 1.0);initialEpsilon:初始探索概率;annealInterval:探索概率衰减所需的步数区间;minEpsilon:epsilon 衰减的下限,低于此值不再减小;decayRate(可选):每次权重更新时模型变化的程度,默认 1.0。
从源码可以看到衰减机制:delta = ((initialEpsilon - minEpsilon) * decayRate) / annealInterval,每次调用Anneal()时epsilon -= delta并用minEpsilon兜底;Sample()中当随机数小于 epsilon 且非确定性模式时随机选取动作,否则选择动作价值最大的动作。
对于异步学习,还需要AggregatedPolicy(policy/aggregated_policy.hpp)。它内部持有多个子策略,每个时间步按给定概率分布随机挑选一个子策略来采样动作。其构造函数要求概率分布向量(arma::colvec)的元素个数与子策略数量一致,且所有元素之和等于 1。
回放(Replay)
最简单的回放是随机回放(Random Replay):每个时间步,智能体与环境交互的经验被保存到内存缓冲区,训练时从缓冲区随机采样过往经验。
mlpack 提供两种回放实现:
RandomReplay(replay/random_replay.hpp):FIFO 环形缓冲区,构造函数为RandomReplay(batchSize, capacity, nSteps = 1, dimension = StateType::dimension),其中batchSize是每次采样返回的样本数,capacity是内存中存储的最大样本数,nSteps是前瞻步数(用于 n-step 学习),dimension是编码状态的维度。PrioritizedReplay(replay/prioritized_replay.hpp):基于 SumTree 数据结构的优先级回放,更多采样高优先级(高 TD 误差)的转移样本,使学习更高效。构造函数为PrioritizedReplay(batchSize, capacity, alpha, nSteps = 1, dimension = ...),其中alpha表示优先级化程度(0 表示纯均匀采样,越大越偏向高优先级样本)。源码中还内置了重要性采样权重 beta(初始 0.6)随训练进度线性退火,以校正优先级采样带来的偏差。
实例化示例
以CartPole环境为例,创建 Greedy Policy 与 Prioritized Replay:
GreedyPolicy<CartPole> policy(1.0, 1000, 0.1); PrioritizedReplay<CartPole> replayMethod(10, 10000, 0.6);policy的三个参数分别是:初始 epsilon 值(1.0)、epsilon 衰减的步数区间(1000)、epsilon 的衰减下限(0.1);replayMethod的三个参数分别是:每次采样返回的批次大小(10)、内存中存储的样本数(10000)、优先级化程度(0.6)。
TrainingConfig:智能体超参数配置
除策略与回放外,RL 智能体还需要大量训练超参数,从未来奖励的折扣率到是否使用 Double Q-Learning,都可通过TrainingConfig类统一配置。该类定义在 training_config.hpp,采用“取值器/修改器”风格接口(如config.StepSize()既可作为 getter 也可作为 setter)。示例配置如下:
TrainingConfig config; config.StepSize() = 0.01; config.Discount() = 0.9; config.TargetNetworkSyncInterval() = 100; config.ExplorationSteps() = 100; config.DoubleQLearning() = false; config.StepLimit() = 200;该配置描述了一个:优化步长(学习率)为 0.01、折扣因子为 0.9、目标网络同步间隔为 100、先积累 100 步探索样本再开始学习、单回合最多 200 步、不使用 Double Q-Learning 的智能体。
结合源码中的构造函数默认值,TrainingConfig的全部可配置项如下表:
| 配置项 | 默认值 | 适用智能体 | 含义 |
|---|---|---|---|
NumWorkers() | 1 | 异步学习 | 并行 worker 数量 |
UpdateInterval() | 1 | 异步学习 | 更新间隔,类似批大小但逐条更新 |
TargetNetworkSyncInterval() | 100 | Q-Learning / 异步 | 目标网络同步间隔 |
StepLimit() | 200 | Q-Learning / 异步 | 每回合最大步数,0 表示无限制 |
ExplorationSteps() | 1 | Q-Learning | 学习开始前积累的探索步数 |
StepSize() | 0.01 | 全部 | 优化器步长(学习率) |
Discount() | 0.99 | 全部 | 未来奖励折扣率 |
GradientLimit() | 40 | 异步学习 | 梯度裁剪上限 |
DoubleQLearning() | false | Q-Learning | 是否启用 Double Q-Learning |
NoisyQLearning() | false | Q-Learning | 是否启用 Noisy Q-Learning |
IsCategorical() | false | Q-Learning | 是否使用 Categorical(分布式)Q 网络 |
AtomSize() | 51 | Categorical | 价值分布的原子(bin)数量 |
VMin()/VMax() | 0 / 200 | Categorical | 价值分布支持的上下界 |
Rho() | 0.005 | SAC | 软更新 Q 网络的系数 |
需要特别说明的是,主教程文档中的示例将TargetNetworkSyncInterval描述为“sync interval of 200 episodes”,但从源码注释看该参数的实际语义是同步间隔(步数维度),实际使用时建议以代码注释和训练行为为准。
五大强化学习智能体
mlpack 提供了多个功能强大的强化学习智能体,覆盖离散动作与连续动作两大场景。每个智能体都高度可定制、可扩展,可灵活应用于各类环境与任务。
Q-Learning(DQN / Double DQN)
Q-Learning 是 mlpack 中最基础的深度强化学习智能体,类定义位于 q_learning.hpp,实现了 DQN、Double DQN 等算法,模板参数为QLearning<EnvironmentType, NetworkType, UpdaterType, PolicyType, ReplayType>。完整的分步讲解与可运行示例见 Q-Learning 教程,核心步骤如下:
- 构建网络。最简单的方式是使用
SimpleDQN(q_networks/simple_dqn.hpp),它创建一个带两个隐藏层的前馈网络:
SimpleDQN<> model(4, 64, 32, 2);输入维度 4 对应CartPole状态的四个成员(位置、速度、角度、角速度),输出维度 2 对应两个动作(forward与backward)。
也可以使用 mlpack 的 ann 模块构建自定义FFN网络,再包装进SimpleDQN(因为 Q-Learning 智能体要求网络对象具备ResetNoise方法,FFN本身没有):
FFN<MeanSquaredError, GaussianInitialization> network(MeanSquaredError(), GaussianInitialization(0, 0.001)); network.Add<Linear>(128); network.Add<ReLU>(); network.Add<Linear>(128); network.Add<ReLU>(); network.Add<Linear>(2); SimpleDQN<> model(network);- 配置策略、回放与超参数:
GreedyPolicy<CartPole> policy(1.0, 1000, 0.1, 0.99); RandomReplay<CartPole> replayMethod(10, 10000); TrainingConfig config; config.StepSize() = 0.01; config.Discount() = 0.9; config.TargetNetworkSyncInterval() = 100; config.ExplorationSteps() = 100; config.DoubleQLearning() = false; config.StepLimit() = 200;- 声明智能体并训练。使用
decltype(var)作为冗长模板类型的简写,传入各对象引用:
QLearning<CartPole, decltype(model), AdamUpdate, decltype(policy)> agent(config, model, policy, replayMethod);训练循环中,每次调用agent.Episode()执行一个回合,用arma::running_stat<double>计算平均回报;当平均回报超过预设阈值(如 35)即认为收敛;若 1000 个回合仍未达标则判定失败。从 q_learning_impl.hpp 可以看到,QLearning构造时会深拷贝一份网络作为targetNetwork,并根据环境的InitialSample()维度重置网络输入尺寸——这也是环境接口必须完整实现的原因。
Asynchronous Learning(异步学习)
2016 年 DeepMind 与蒙特利尔大学的研究者发表了论文Asynchronous Methods for Deep Reinforcement Learning,提出四种异步算法:One-Step SARSA、One-Step Q-Learning、N-Step Q-Learning 与 Advantage Actor-Critic(A3C)。
在线 RL 算法与深度神经网络的组合因在线更新的非平稳性和相关性而很不稳定;经验回放虽然解决了这一问题,但占用更多内存与算力,且要求 off-policy 算法。异步方法不再使用经验回放,而是让多个智能体在多个环境实例上并行异步执行,一举解决上述问题。详细教程见 异步学习教程。
mlpack 的异步学习把每个并行智能体称为“worker”,当前实现包括 One-Step Q-Learning worker、N-Step Q-Learning worker 与 One-Step SARSA worker(见 worker/)。不使用单一策略,而是用AggregatedPolicy聚合多个子策略,每个子策略对应一个 worker,worker 数量取决于子策略数量:
AggregatedPolicy<GreedyPolicy<CartPole>> policy({GreedyPolicy<CartPole>(0.7, 5000, 0.1), GreedyPolicy<CartPole>(0.7, 5000, 0.01), GreedyPolicy<CartPole>(0.7, 5000, 0.5)}, arma::colvec("0.4 0.3 0.3"));概率分布向量"0.4 0.3 0.3"的元素个数必须与子策略数量一致,且元素之和等于 1。随后声明 One-Step Q-Learning 智能体(也可按需改用NStepQLearning或OneStepSarsa):
OneStepQLearning<CartPole, decltype(model), ens::AdamUpdate, decltype(policy)> agent(std::move(config), std::move(model), std::move(policy));与 Q-Learning 不同,异步学习使用Train()方法训练,并传入一个返回布尔值的measure回调(lambda),它接收确定性测试回合的总回报作为参数,返回 true 表示训练结束:
for (int i = 0; i < 100; i++) { agent.Train(measure); }measure的典型实现会维护一个固定长度的回报环形缓冲,打印每回合回报与平均回报,并在超过最大回合数时返回 true。这样三个智能体就在三个 CPU 线程上异步训练,共同更新动作价值估计。从 async_learning.hpp 的注释可以看到,Measure需要是形如bool foo(double reward)的可调用对象。
Deep Deterministic Policy Gradient(DDPG)
DDPG 结合了 Q-Learning 与策略梯度的优势,非常适合连续动作空间问题,通过深度神经网络同时逼近 Q 值与策略函数,学习从状态直接映射到动作的确定性策略。详细教程见 DDPG 教程,类定义位于 ddpg.hpp。
以Pendulum环境为例,先配置回放与训练参数:
RandomReplay<Pendulum> replayMethod(32, 10000); TrainingConfig config; config.StepSize() = 0.01; config.TargetNetworkSyncInterval() = 1; config.UpdateInterval() = 3;接着构建 Actor(生成动作)与 Critic(评估动作价值)两个网络。Actor 输出层使用TanH将动作约束在 [-1, 1]:
FFN<EmptyLoss, GaussianInitialization> policyNetwork(EmptyLoss(), GaussianInitialization(0, 0.1)); policyNetwork.Add(new Linear(128)); policyNetwork.Add(new ReLU()); policyNetwork.Add(new Linear(1)); policyNetwork.Add(new TanH()); FFN<EmptyLoss, GaussianInitialization> qNetwork(EmptyLoss(), GaussianInitialization(0, 0.1)); qNetwork.Add(new Linear(128)); qNetwork.Add(new ReLU()); qNetwork.Add(new Linear(1));噪声是 DDPG 探索的关键。mlpack 提供两种噪声:
GaussianNoise(noise/gaussian.hpp):高斯分布不相关噪声,参数为(size, mu, sigma);OUNoise(noise/ornstein_uhlenbeck.hpp):Ornstein-Uhlenbeck 过程产生时间相关的噪声,适合有惯性/动量的物理控制问题,参数为(size, mu, theta, sigma),其中theta是均值回归速率,sigma是噪声标准差,默认值分别为 0.15 与 0.2。
OUNoise ouNoise(size, mu, theta, sigma); DDPG<Pendulum, decltype(qNetwork), decltype(policyNetwork), OUNoise, AdamUpdate> agent(config, qNetwork, policyNetwork, ouNoise, replayMethod);训练与收敛判定:设置回报阈值(如 -400)、最大回合数(如 1000)与连续测试回合数(如 10),维护最近若干回合的平均回报;当连续 10 个回合平均回报超过阈值时,置agent.Deterministic() = true执行 10 个确定性测试回合验证收敛效果。
Twin Delayed Deep Deterministic Policy Gradient(TD3)
TD3 在 DDPG 基础上引入若干创新以提升训练稳定性:使用孪生 Q 网络(twin Q-networks)缓解传统 Q-Learning 固有的过估计偏差,并采用目标策略平滑技术提升策略更新的稳定性。详细教程见 TD3 教程,类定义位于 td3.hpp。
与 DDPG 的配置差异在于:TD3不需要噪声类型,并在配置中多出一个Rho()参数用于目标网络的软更新:
RandomReplay<Pendulum> replayMethod(32, 10000); TrainingConfig config; config.StepSize() = 0.01; config.TargetNetworkSyncInterval() = 1; config.UpdateInterval() = 3; config.Rho() = 0.001;Actor 与 Critic 网络的构建方式与 DDPG 一致,随后实例化智能体:
TD3<Pendulum, decltype(qNetwork), decltype(policyNetwork), AdamUpdate> agent(config, qNetwork, policyNetwork, replayMethod);收敛判定的训练循环与 DDPG 示例结构相同,教程中Pendulum环境的回报阈值设为 -300。
Soft Actor-Critic(SAC)
SAC 是面向连续动作空间的进阶算法,采用软 Q 网络缓解过估计偏差,并通过熵最大化鼓励探索,从而获得更鲁棒、稳定的策略学习。详细教程见 SAC 教程,类定义位于 sac.hpp。
配置与 TD3 相同,同样使用Rho()控制 Q 网络的软更新:
RandomReplay<Pendulum> replayMethod(32, 10000); TrainingConfig config; config.StepSize() = 0.01; config.TargetNetworkSyncInterval() = 1; config.UpdateInterval() = 3; config.Rho() = 0.001;Actor(策略)与 Critic(价值)网络的构建与 DDPG/TD3 一致(Actor 输出层同样使用TanH),随后实例化:
SAC<Pendulum, decltype(qNetwork), decltype(policyNetwork), AdamUpdate> agent(config, qNetwork, policyNetwork, replayMethod);SAC 的训练与收敛判定循环与 TD3 示例一致(回报阈值 -300)。从 training_config.hpp 的注释可知,Rho()是 SAC 专属的软更新 Q 网络参数。
深入源码与进一步阅读
mlpack 强化学习框架的全部实现位于 src/mlpack/methods/reinforcement_learning/,包含环境(environment/)、策略(policy/)、回放(replay/)、Q 网络(q_networks/)、worker(worker/)、噪声(noise/)与智能体类,每份头文件都带有详尽的 Doxygen 注释,是深入了解各算法实现细节的第一手资料。
对应的测试代码同样值得阅读:其中 q_learning_test.cpp 与 policy_gradient_test.cpp 覆盖了 Q-Learning、异步学习等智能体的训练收敛验证,可作为“如何正确配置与断言训练效果”的参考范例。
各智能体的分步教程与完整可运行代码位于同一目录下:
- Q-Learning 教程
- 异步学习教程
- DDPG 教程
- TD3 教程
- SAC 教程
其中每一篇都包含可直接在本地编译运行的完整程序。你可以在此基础上自由调整网络结构、策略参数、回放容量与各超参数,观察它们对训练过程与最终收敛效果的影响——这正是在 mlpack 中深入掌握强化学习的最佳方式。
- 人工智能
- 机器学习
- 深度学习
【免费下载链接】mlpack
mlpack: a fast, header-only C++ machine learning library
相关推荐
如何扩展HighlightedTextEditor功能:创建自定义预设与正则表达式库的最佳实践 🚀
如何扩展HighlightedTextEditor功能:创建自定义预设与正则表达式库的最佳实践 🚀 SwiftUI的HighlightedTextEditor
darkhttpd源码解析:从事件循环到HTTP请求处理的核心实现
darkhttpd源码解析:从事件循环到HTTP请求处理的核心实现 darkhttpd是一款轻量级单线程静态Web服务器,以简洁高效著称。本文将深入剖析其核心实
PyMARL:多智能体强化学习的模块化实战框架
PyMARL:多智能体强化学习的模块化实战框架 在人工智能技术飞速发展的今天,多智能体强化学习(MARL)正成为解决复杂协同决策问题的关键技术。PyMARL作为
人工智能强化学习多智能体深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考