强化学习进阶篇八 策略梯度算法

简介

Q-learning、DQN 及 DQN 改进算法都是基于价值 的方法。其中 Q-learning 是处理有限状态 的算法,而 DQN 可以用来解决连续状态 的问题。除了基于值函数的方法,还有一支非常经典的方法,那就是基于策略的方法。在此之前要学一些基础:

  • 深度学习-交叉熵 最大似然估计 MLE

  • 梯度

    f ( x ) = x 2 + 2 f(x)=x^2+2 f(x)=x2+2 这是开口向上的抛物线,

    1. 一阶导数 (梯度 ): f ′ ( x ) = 2 x f'(x)=2x f′(x)=2x 变化率、切线斜率

      • x < 0 , x 递增 x<0,x递增 x<0,x递增 f'(x)=2x \<0 也是在递增,这叫梯度上升。原 f ( x ) = x 2 + 2 f(x)=x^2+2 f(x)=x2+2是递减
      • 0 < x , x 递时 0<x,x递时 0<x,x递时 f'(x)=2x \>0 也是在递增,原 f ( x ) = x 2 + 2 f(x)=x^2+2 f(x)=x2+2是递增
    2. 二阶导数: f ′ ′ ( x ) = 2 f''(x)=2 f′′(x)=2 加速度信息

      • f ′ ′ ( x ) = 2 > 0 f''(x)=2>0 f′′(x)=2>0 极小值点
      • 若 f ′ ′ ( x ) < 0 f''(x)<0 f′′(x)<0 极大值点

策略梯度原理

公式推导过程

MLE公式 ℓ ( θ ) = − 1 m ∑ i = 1 m ln ⁡ p ( x i ; θ ) \ell(\theta) = - \frac{1}{m} \sum_{i=1}^m \ln p(x_i;\theta) ℓ(θ)=−m1∑i=1mlnp(xi;θ) 在抛硬币例子中 寻找参数 θ \boldsymbol{\theta} θ ,根据结果反推原因的核心思想。

策略梯度的核心思想和MLE公式相同,寻找最优参数 θ \boldsymbol{\theta} θ。 先回顾一下:

状态价值函数 V π ( s ) = E π G t ∣ S t = s = ∑ a ′ ∈ A π ( s , a ) Q π ( s , a ) V^{\pi}(s)=\mathbb{E}{\pi}G_t\|S_t=s= \sum{a'\in \mathcal{A}} \pi(s,a) Q^{\pi}(s,a) Vπ(s)=EπGt∣St=s=∑a′∈Aπ(s,a)Qπ(s,a)

动作价值函数 Q π ( s , a ) = E π G t ∣ A t = a , S t = s = r ( s , a ) + γ ∑ p ( s ′ ∣ s , a ) V π ( s ′ ) Q^{\pi}(s,a)= \mathbb{E}_{\pi}G_t\|A_t=a,S_t=s=r(s,a)+\gamma \sum p(s'|s,a) V^{\pi}(s') Qπ(s,a)=EπGt∣At=a,St=s=r(s,a)+γ∑p(s′∣s,a)Vπ(s′)

定义目标函数 J ( θ ) = E π V π ( s ) = ∑ a ′ ∈ A π ( s , a ) ⏟ θ r ( s , a ) + γ ∑ p ( s ′ ∣ s , a ) V π ( s ′ ) ⏟ G t = Q π ( s , a ) J(\theta)=\mathbb{E}{\pi}V\^{\\pi}(s)=\sum{a'\in \mathcal{A}} \underbrace{\pi(s,a)}\theta \underbrace{r(s,a)+\\gamma \\sum p(s'\|s,a) V\^{\\pi}(s')}{G_t=Q^{\pi}(s,a)} J(θ)=EπVπ(s)=∑a′∈Aθ π(s,a)Gt=Qπ(s,a) r(s,a)+γ∑p(s′∣s,a)Vπ(s′)

简化一下公式: J ( θ ) = ∑ π ( s , a ) ∗ Q π ( s , a ) J(\theta)=\sum \pi(s,a)*Q^{\pi}(s,a) J(θ)=∑π(s,a)∗Qπ(s,a) 其中 π ( s , a ) \pi(s,a) π(s,a)是概率函数: p ( x ; θ ) = π ( s , a ) p(x;\theta)=\pi(s,a) p(x;θ)=π(s,a), Q π ( s , a ) Q^{\pi}(s,a) Qπ(s,a) 表示权重。

策略函数建模: π θ ( a ∣ s ) ⏟ 动作概率分布 = s o f t m a x ( w 2   max ⁡ ( w 1 s + b 1 ,   0 ) + b 2 ) ⏟ 策略神经网络 , θ = { w 1 , b 1 , w 2 , b 2 } ⏟ 网络参数 \underbrace{\pi_{\boldsymbol{\theta}}(a|s)}{动作概率分布}= \underbrace{\mathrm{softmax}\Big(w_2\, \max\big(w_1 s + b_1,\,0\big)+b_2\Big)}{策略神经网络},\quad \boldsymbol{\theta}= \underbrace {\{w_1,b_1,w_2,b_2\}}_{网络参数} 动作概率分布 πθ(a∣s)=策略神经网络 softmax(w2max(w1s+b1,0)+b2),θ=网络参数 {w1,b1,w2,b2}

作为策略学习目标,寻找最优参数 θ \boldsymbol{\theta} θ。网络输入状态 s s s,输出每个动作的概率。例如车杆环境有向左和向右两个动作,策略网络可能输出:

π θ ( ⋅ ∣ s ) = 0.3 , 0.7 \pi_\theta(\cdot|s)=0.3,0.7 πθ(⋅∣s)=0.3,0.7

表示当前状态下:

  • 选择动作 a 1 a_1 a1 的概率为 0.3 0.3 0.3;
  • 选择动作 a 2 a_2 a2 的概率为 0.7 0.7 0.7。

J ( θ ) J(\theta) J(θ)目标函数用对数似然函数 J ( θ ) = ∑ i = 1 m G i ∗ ln ⁡ π θ ( a i ∣ s ) J(\theta)=\sum_{i=1}^m G_i * \ln \pi_\theta(a_i|s) J(θ)=∑i=1mGi∗lnπθ(ai∣s)

最后对 J ( θ ) J(\theta) J(θ)求导 ∇ J ( θ ) = ∑ i = 1 m G i ∗ ∇ ln ⁡ π θ ( a i ∣ s ) \nabla J(\theta)= \sum_{i=1}^m G_i* \nabla \ln \pi_\theta(a_i|s) ∇J(θ)=∑i=1mGi∗∇lnπθ(ai∣s)

策略梯度公式 ∇ J ( θ ) = ∑ i = 1 m G i ∗ ∇ ln ⁡ π θ ( a i ∣ s ) \nabla J(\theta)= \sum_{i=1}^m G_i* \nabla \ln \pi_\theta(a_i|s) ∇J(θ)=∑i=1mGi∗∇lnπθ(ai∣s)

策略梯度计算演示

动作空间: a 1 , a 2 a_1,a_2 a1,a2,状态空间: s 0 , s 1 , s 2 s_0,s_1,s_2 s0,s1,s2

采样轨迹: ( s 0 , a 1 , r = 1 ) , ( s 1 , a 2 , r = 1 ) , ( s 2 , a 2 , r = 0 ) (s_0,a_1,r=1),(s_1,a_2,r=1),(s_2,a_2,r=0) (s0,a1,r=1),(s1,a2,r=1),(s2,a2,r=0)

折扣系数 γ = 1 \gamma=1 γ=1

回合长度 ( T = 3 ) (T=3) (T=3),时间下标 ( t = 0 , 1 , 2 ) (t=0,1,2) (t=0,1,2)

回报 G t G_t Gt汇总: G t = ∑ k = 0 ∞ γ k R t + k = > ∑ k = 0 R t + k G_t=\sum_{k=0}^{\infty}\gamma^{k} R_{t+k}=>\sum_{k=0} R_{t+k} Gt=∑k=0∞γkRt+k=>∑k=0Rt+k

  • G 2 = r 2 = 0 G_2=r_2=0 G2=r2=0

  • G 1 = r 1 + G 2 = 1 + 0 = 1 G_1=r_1+G_2=1+0=1 G1=r1+G2=1+0=1

  • G 0 = r 0 + G 1 + G 2 = 1 + 1 + 0 = 2 G_0=r_0+G_1+G_2=1+1+0=2 G0=r0+G1+G2=1+1+0=2

汇总表

t t t s t s_t st a t a_t at r t r_t rt G t G_t Gt
0 s 0 s_0 s0 a 1 a_1 a1 1 2
1 s 1 s_1 s1 a 2 a_2 a2 1 1
2 s 2 s_2 s2 a 2 a_2 a2 0 0

根据策略梯度公式 ∇ J ( θ ) = ∑ i = 1 m G i ∗ ∇ ln ⁡ π θ ( a i ∣ s ) \nabla J(\theta)= \sum_{i=1}^m G_i* \nabla \ln \pi_\theta(a_i|s) ∇J(θ)=∑i=1mGi∗∇lnπθ(ai∣s)

策略梯度计算过程

  1. t = 0 t=0 t=0, G 0 = 2 G_0=2 G0=2, s 0 , a 1 s_0, a_1 s0,a1: 2 ∗ ∇ ln ⁡ π θ ( a 1 ∣ s 0 ) 2*\nabla \ln \pi_\theta(a_1|s_0) 2∗∇lnπθ(a1∣s0)

  2. t = 1 t=1 t=1, G 1 = 1 G_1=1 G1=1, s 1 , a 2 s_1, a_2 s1,a2: 1 ∗ ∇ ln ⁡ π θ ( a 2 ∣ s 1 ) 1*\nabla \ln \pi_\theta(a_2|s_1) 1∗∇lnπθ(a2∣s1)

  3. t = 2 t=2 t=2, G 2 = 0 G_2=0 G2=0, s 2 , a 2 s_2, a_2 s2,a2: 0 ∗ ∇ ln ⁡ π θ ( a 2 ∣ s 2 ) 0*\nabla \ln \pi_\theta(a_2|s_2) 0∗∇lnπθ(a2∣s2)

  4. 根据策略梯度公式 ∇ J ( θ ) = 2 ∗ ∇ ln ⁡ π θ ( a 1 ∣ s 0 ) + 1 ∗ ∇ ln ⁡ π θ ( a 2 ∣ s 1 ) + 0 \nabla J(\theta)= 2\*\\nabla \\ln \\pi_\\theta(a_1\|s_0) + 1\*\\nabla \\ln \\pi_\\theta(a_2\|s_1) + 0 ∇J(θ)=2∗∇lnπθ(a1∣s0)+1∗∇lnπθ(a2∣s1)+0,最大化期望回报 J ( θ ) J(\theta) J(θ),要用梯度上升 ,深度学习默认做梯度下降 所以乘以负号 ∇ J ( θ ) = − 2 ∗ ∇ ln ⁡ π θ ( a 1 ∣ s 0 ) + 1 ∗ ∇ ln ⁡ π θ ( a 2 ∣ s 1 ) + 0 \nabla J(\theta)= - 2\*\\nabla \\ln \\pi_\\theta(a_1\|s_0) + 1\*\\nabla \\ln \\pi_\\theta(a_2\|s_1) + 0 ∇J(θ)=−2∗∇lnπθ(a1∣s0)+1∗∇lnπθ(a2∣s1)+0

实现关键代码

c++ 复制代码
计算轨迹策略梯度
 int length = vList.size()-1;
 double G = 0;
 adam.zero_grad();

 for (int i = length; 0 <= i; i--)
 {
     // s,a,r,
     auto [s, a, r, s1, d] = vList[i];
     auto s0 = VectorDoubleTensor(s);
     auto act = torch::tensor({ {a} }, torch::kInt);

     auto action = m_Qnet->forward(s0).gather(1, act);
     auto logprob = torch::log(action + 1e-8);
     G = m_dbGamma * G + r;
     auto lass = -logprob * G;
     lass.backward();

 }
 adam.step();

策略梯度实现细节

0. 策略网络 π θ \pi_\theta πθ

复制代码
class PolicyNetImpl : public torch::nn::Module
{
public:
    PolicyNetImpl() = default;

    PolicyNetImpl(int64_t input, int64_t output, int64_t hidden = 128)
    {
        m_fc1 = register_module("fc1", torch::nn::Linear(input, hidden));
        m_fc2 = register_module("fc2", torch::nn::Linear(hidden, output));
    }

    torch::Tensor forward(torch::Tensor x)
    {
        x = torch::relu(m_fc1->forward(x));
        x = m_fc2->forward(x);
        return torch::softmax(x, 1);
    }

    torch::nn::Linear m_fc1{ nullptr };
    torch::nn::Linear m_fc2{ nullptr };
};

TORCH_MODULE(PolicyNet);

网络结构包括一个隐藏层和一个输出层:

  1. 输入是环境状态 s s s;
  2. 隐藏层使用 ReLU 激活函数;
  3. 输出层维度等于动作数量;
  4. softmax 将网络输出转换成概率分布,并保证所有动作概率之和为 1 1 1。

1. 创建策略网络和优化器

复制代码
    auto input = m_objEnv->GetStateDim();
    auto output = m_objEnv->GetActionDim();

    m_Qnet = PolicyNet(input, output);
    m_Qnet->to(m_device);

    CreateOptimizer(m_Qnet);

2. 根据策略选择动作

复制代码
double PolicyGradient::TakeAction(VectorDouble& s0, bool bPredict)
{
    torch::NoGradGuard no_grad;
    auto s = VectorDoubleTensor(s0, m_device);
    auto logits = m_Qnet->forward(s);
    torch::Tensor action;

    if (bPredict)
    {
        action = logits.argmax(-1);
    }
    else
    {
        Categorical categorical(logits);
        action = categorical.sample();
    }

    return action.item<int>();
}

3. 收集完整回合数据

蒙特卡洛算法 采集完整轨迹数据,通过 TrainGenerateItem2() 使用整条轨迹更新策略,从后向前计算累计回报。

复制代码
void PolicyGradient::TrainGenerateItem2(const QwList& vList)
{
    if (vList.size() == 0)
    {
        return;
    }

    if (450 < vList.size())
    {
        m_bEndGenerateTrain = true;
        return;
    }

    int length = vList.size() - 1;
    double G = 0;
    m_pAdam->zero_grad();

    for (int i = length; 0 <= i; i--)
    {
        // s,a,r,
        auto [s, a, r, s1, d] = vList[i];
        auto s0 = VectorDoubleTensor(s, m_device);
        auto act = torch::tensor({ {a} }, torch::kInt).to(m_device);

        auto action = m_Qnet->forward(s0).gather(1, act);
        auto logprob = torch::log(action + 1e-8);
        G = m_dbGamma * G + r;
        auto lass = -logprob * G;
        lass.backward();

    }
    m_pAdam->step();
}

4. 运行效果

复制代码
int main()
{
    DeepQNetwork  deepQN;
    //deepQN.Play(400);
    //deepQN.DoubleDQN(400);
    DuelingDQN duelingDQN;
    //duelingDQN.Play(400);

    PolicyGradient policyGradient;
    policyGradient.Play(1000);
    /**
    Currently PolicyGradient
GenerateTrainData .....
train i: 80 / 1000 , rewardCount: 36
train i: 100 / 1000 , rewardCount: 44
train i: 120 / 1000 , rewardCount: 43
train i: 140 / 1000 , rewardCount: 184
train i: 160 / 1000 , rewardCount: 121
train i: 180 / 1000 , rewardCount: 17
train i: 200 / 1000 , rewardCount: 61
train i: 220 / 1000 , rewardCount: 132
train i: 240 / 1000 , rewardCount: 331
train i: 260 / 1000 , rewardCount: 55
train i: 280 / 1000 , rewardCount: 94
train i: 300 / 1000 , rewardCount: 245
train i: 319, break Generate Train ####

TestData .....
count: 1 , rewardCount: 500
count: 2 , rewardCount: 456
count: 3 , rewardCount: 420
count: 4 , rewardCount: 445
count: 5 , rewardCount: 500
count: 6 , rewardCount: 449
count: 7 , rewardCount: 438
count: 8 , rewardCount: 500
count: 9 , rewardCount: 500
count: 10 , rewardCount: 332
count: 11 , rewardCount: 500
count: 12 , rewardCount: 339
count: 13 , rewardCount: 373
count: 14 , rewardCount: 500
count: 15 , rewardCount: 500
count: 16 , rewardCount: 401
count: 17 , rewardCount: 341
count: 18 , rewardCount: 500
count: 19 , rewardCount: 374
count: 20 , rewardCount: 500
count: 21 , rewardCount: 259
count: 22 , rewardCount: 356
count: 23 , rewardCount: 500
count: 24 , rewardCount: 441
count: 25 , rewardCount: 397
count: 26 , rewardCount: 500
count: 27 , rewardCount: 500
count: 28 , rewardCount: 500
count: 29 , rewardCount: 500
count: 30 , rewardCount: 489
count: 31 , rewardCount: 500
count: 32 , rewardCount: 408
count: 33 , rewardCount: 500
count: 34 , rewardCount: 500
count: 35 , rewardCount: 427
count: 36 , rewardCount: 381
count: 37 , rewardCount: 356
count: 38 , rewardCount: 336
count: 39 , rewardCount: 500
count: 40 , rewardCount: 427
count: 41 , rewardCount: 500
count: 42 , rewardCount: 500
count: 43 , rewardCount: 500
count: 44 , rewardCount: 365
count: 45 , rewardCount: 500
count: 46 , rewardCount: 385
count: 47 , rewardCount: 500
count: 48 , rewardCount: 500
count: 49 , rewardCount: 500
count: 50 , rewardCount: 464
count: 51 , rewardCount: 500
count: 52 , rewardCount: 274
count: 53 , rewardCount: 500
count: 54 , rewardCount: 390
count: 55 , rewardCount: 325
count: 56 , rewardCount: 500
count: 57 , rewardCount: 500
count: 58 , rewardCount: 500
count: 59 , rewardCount: 500
count: 60 , rewardCount: 500
count: 61 , rewardCount: 419
count: 62 , rewardCount: 500
count: 63 , rewardCount: 432
count: 64 , rewardCount: 403
count: 65 , rewardCount: 500
count: 66 , rewardCount: 500
    **/

    }
相关推荐
乌萨奇也要立志学C++1 小时前
【洛谷】kmp算法
开发语言·算法
禹凕1 小时前
Dijkstra算法详解与应用
python·算法
star learning white1 小时前
STM32入门学习5
stm32·嵌入式硬件·学习
传奇开心果编程1 小时前
【Rust入门知识点学与练】第4课:条件判断 if / else
开发语言·学习·rust
lifallen2 小时前
Skill 的生命周期:问题消失以后
人工智能·学习·ai·ai编程
千谦阙听2 小时前
C++类和对象(中):默认成员函数、构造与析构、拷贝构造、运算符重载
开发语言·c++·学习
YaraMemo2 小时前
元启发式算法框架
人工智能·算法·5g·信息与通信·启发式算法·信号处理
weixin_307779132 小时前
一维无粘 Burgers 方程的激波形成问题:MacCormack 格式求解
c++·算法·matlab
Agudamu11613 小时前
AI 视频学习笔记工具拆解:4款工具谁能把视频变成复习资料
人工智能·学习·音视频