[学习笔记]trpo——对策略进行显式约束

再继续,actor-critic之后就是著名的trpo

这个东西熬,算是强化学习入门之后的第一个boss了

第一遍看完,只觉得它是策略梯度的pro plus版本,后续看来,它是能作为接下来好几年开山之作的存在

0.Actor-Critic算法的优劣分析

首先,还是分析一下之前Actor-Critic算法的优劣吧

有一个从机器学习入门就困扰大家的问题:步长

步子太大,容易扯着蛋,步子太小,又难以收敛

在之前的reinforce之中,就是因为步长可能会太大,会导致收敛不稳定,两次训练的结果可能一次收敛一次不收敛

而Actor-Critic其实也没有解决这个问题,只是Actor-Critic只是解决了“判标”问题

实际上,这牵扯到一个更加关键的问题:我们并没有把策略真正地量化表达出来

回想一下,从reinforce开始,我们只是借助了神经网络的超强拟合能力,强行表达策略

但是我们真的不知道策略,到底是什么,不是么?

而trpo,就在尝试解决这个问题,让训练变得稳定可靠,用数学的办法。

1.trpo解决问题的办法

首先,需要解决的问题,就是训练的稳定性。

在强化学习中,步长代表了什么呢?

通俗来说,步长就是两种策略之间的差别。

比如,上上下下左左右右BABA和上上下下左左右右ABAB,就是间隔步长比较小的策略,而ABAB右右左左下下上上,就和第一种策略不一样。

那么,如何“量化”策略之间的差别呢?

直接遍历?那时间复杂度会指数级增长,显然不可能。

一个直观的想法,就是对每个策略都进行单位向量化,然后通过内积等方法来进行判断

不过呢,这个方法显然太初级了,还得是大佬的方法直截了当

尽管我尽全力在避免数学,但是这里还是不得不插入一条公式:

{D_{KL}(\pi_{\theta_{\text{old}}}(\cdot|s) \| \pi_{\theta}(\cdot|s))}

说实话,第一次看到这条公式,我就立马想到了另外一个东西:信息熵:

{H(\pi(\cdot|s)) = -\sum_{a} \pi(a|s) \log \pi(a|s)}

怎么样,是不是很像很像?

反正我是没有找到网上任何一个对于kl散度和信息熵的联系的讲解,我觉得从这方面来讲就非常非常非常的直观

信息熵,如果系统学习过机器学习,那就会了解到,它是计算loss的一个过程。

信息熵衡量了一个过程的“平均惊奇程度”,不了解的话可以去看一下,还是挺简单的

那么,kl散度,不就是两个状态信息熵的差值吗??!!

那kl散度是不是就是两个策略的“平均惊奇程度”的差值?

也就是两个策略的“相对惊奇程度”?

这个工具,就可以用来衡量两种策略的差别了,只要不超过阈值,就可以

放 心 大 胆 地 放 大 步 长 !

(这里插一嘴,因为要对比策略的相似性,所以我们就不能够使用经验回放了。经验回放是什么,是将策略给存储下来,并在要用的时候随机拿出来。这样就会不可避免地会取出很多“毫不相关”的策略,这不符合我们使用kl散度的目的,所以还是放弃了吧)

接下来,解决了步长问题,就该是方向问题了。

在之前地策略更新时,我们使用的梯度方向都是loss的方向,也就是原始网络的方向

 但是,在配合上由kl散度确定的最大步长后,这个步长下的更新很可能是非常危险的,总不能计算了步长之后,再回头判断一下策略是否危险,直到策略不危险再更新吧,这也太蠢了。

(实际上,trpo就是这么干的。这一步在trpo中叫做线搜索,就是一个硬约束,如果步长导致不稳定,则减小步长,再更新。而这样会造成更新的速度变慢,这个问题的解决要到ppo来讲了,我们过几天再鸽)

于是乎,需要利用kl散度,来确定一个新的变量空间进行优化。

当然,这一部分也比较直观,用本科学过的最优化的知识就可以解决。

记得最优化里有一堆方法,什么牛顿下山啊最速梯度啊共轭梯度啊啥的

但是这些都离不开一个东西,叫做二阶导数,或者说——hessian矩阵

要对参数进行优化,那必须要计算函数对于每个变量的二阶导

这个东西的复杂度是网络规模的平方,也就是很大很大。

计算它是不可能的。

但是呢,这也是比较巧妙的一点

我们真正需要的不是hessian矩阵,而是矩阵蕴含的“信息”

也就是自然梯度方向,x=F^{−1}∇L

(自然梯度:指策略空间下,在kl散度的度量中计算的最大梯度,跟原始梯度没有关系)

这里也是我没看懂的地方,因为涉及的数学很多,贴一下我看懂的部分吧

首先,优化需要使用KKT条件,这个在最优化中已经很熟悉了

λFΔθ=g

忽略常数λ,我们需要解的就是Δθ=F^{-1}g,自然梯度方向。

不少博客都在说要计算F逆,导致我看得云里雾里的,但其实这不就是简单的解线性方程组么。

设方程组Fx=g,解出来x不就行了

解这个方程组,其实也是比较麻烦的

 

而共轭梯度法可以很好地解决这个问题

共轭梯度法会自动迭代找到这个自然梯度的方向,所以这个问题也就可以解决了

于是,trpo的大概过程就差不多了,理解这个过程也是花了一番功夫

源码:

# TRPO:Actor-Critic强化版,基本思想与AC一致,使用了不同的迭代更新方法。
'''
KL 散度(Kullback-Leibler Divergence)精确地衡量了当我们使用一个近似概率分布 Q 来建模或描述一个真实概率分布 P 时,所引入的信息损失。
简而言之,它量化了“近似”与“真实”之间的差距。KL 散度值越小,意味着分布 Q 对分布 P 的拟合程度越高。
kl散度用来计算策略之间的差别,如果KL散度小于MAX_KL,则策略更新。
如果KL散度大于MAX_KL,则策略更新停止,防止步长过大从而导致扯到蛋。
'''
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import gymnasium as gym
import time
from torch.distributions import Categorical

# Hyper Parameters
BATCH_SIZE = 4000  # TRPO通常使用更大的批量大小
LR = 0.01  # 学习率(主要用于价值网络)
GAMMA = 0.99  # 奖励折扣
LAMBDA = 0.95  # GAE参数
MAX_KL = 0.01  # 最大KL散度
DAMPING = 0.1  # 阻尼系数
env = gym.make('CartPole-v1', render_mode='human')
env = env.unwrapped
N_ACTIONS = env.action_space.n
N_STATES = env.observation_space.shape[0]
ENV_A_SHAPE = 0 if isinstance(
    env.action_space.sample(), int) else env.action_space.sample().shape


class PolicyNet(nn.Module):
    def __init__(self):
        super(PolicyNet, self).__init__()
        self.fc1 = nn.Linear(N_STATES, 50)
        self.fc1.weight.data.normal_(0, 0.1)
        self.out = nn.Linear(50, N_ACTIONS)
        self.out.weight.data.normal_(0, 0.1)

    def forward(self, x):
        x = self.fc1(x)
        x = F.relu(x)
        action_logits = self.out(x)
        return F.softmax(action_logits, dim=-1)


class ValueNet(nn.Module):
    def __init__(self):
        super(ValueNet, self).__init__()
        self.fc1 = nn.Linear(N_STATES, 50)
        self.fc1.weight.data.normal_(0, 0.1)
        self.out = nn.Linear(50, 1)
        self.out.weight.data.normal_(0, 0.1)

    def forward(self, x):
        x = self.fc1(x)
        x = F.relu(x)
        state_value = self.out(x)
        return state_value


class TRPO:
    def __init__(self):
        self.policy_net = PolicyNet()
        self.value_net = ValueNet()
        self.value_optimizer = torch.optim.Adam(
            self.value_net.parameters(), lr=LR)

    def choose_action(self, x):  # 与AC一致,得出概率后采样
        x = torch.unsqueeze(torch.FloatTensor(x), 0)
        probs = self.policy_net(x)
        m = Categorical(probs)
        action = m.sample()
        return action.item(), m.log_prob(action)

    def compute_advantages(self, rewards, values, dones, next_value):
        advantages = np.zeros(len(rewards))
        last_advantage = 0

        for t in reversed(range(len(rewards))):
            if t == len(rewards) - 1:
                next_value = next_value
            else:
                next_value = values[t + 1] * (1 - dones[t])
                # dons:是否为终止态,开关

            delta = rewards[t] + GAMMA * next_value - values[t]
            advantages[t] = delta + GAMMA * LAMBDA * \
                (1 - dones[t]) * last_advantage
            last_advantage = advantages[t]

        return advantages

    def update_value_net(self, states, returns):
        states = torch.FloatTensor(states)
        returns = torch.FloatTensor(returns).unsqueeze(1)

        for _ in range(10):  # 多次更新价值网络
            # 与AC一致,利用均方差损失函数来更新critic网络
            values = self.value_net(states)
            value_loss = F.mse_loss(values, returns)

            self.value_optimizer.zero_grad()
            value_loss.backward()
            self.value_optimizer.step()

    def update_policy(self, states, actions, log_probs_old, advantages):
        states = torch.FloatTensor(states)
        actions = torch.LongTensor(actions)
        log_probs_old = torch.FloatTensor(log_probs_old)
        advantages = torch.FloatTensor(advantages)

        # 计算旧策略的概率
        with torch.no_grad():
            old_probs = self.policy_net(states)
            old_log_probs = torch.log(old_probs.gather(
                1, actions.unsqueeze(1))).squeeze()

        # 计算损失函数的梯度
        def get_loss(volatile=False):
            if volatile:
                with torch.no_grad():
                    probs = self.policy_net(states)
                    log_probs = torch.log(probs.gather(
                        1, actions.unsqueeze(1))).squeeze()
            else:
                probs = self.policy_net(states)
                log_probs = torch.log(probs.gather(
                    1, actions.unsqueeze(1))).squeeze()

            ratio = torch.exp(log_probs - log_probs_old)
            loss = (ratio * advantages).mean()
            return loss

        # 计算KL散度
        def get_kl():
            with torch.no_grad():
                probs = self.policy_net(states)
                log_probs = torch.log(probs)
                old_log_probs_data = torch.log(old_probs)
                kl = (old_probs * (old_log_probs_data - log_probs)
                      ).sum(dim=1).mean()
            return kl

        # 计算梯度
        loss = get_loss()
        grads = torch.autograd.grad(loss, self.policy_net.parameters())
        flat_grad = torch.cat([grad.view(-1) for grad in grads])

        # 计算Fisher-vector乘积
        def fisher_vector_product(v):
            kl = get_kl()
            grads = torch.autograd.grad(
                kl, self.policy_net.parameters(), create_graph=True)
            flat_grad_kl = torch.cat([grad.view(-1) for grad in grads])

            kl_v = (flat_grad_kl * v).sum()
            grads = torch.autograd.grad(kl_v, self.policy_net.parameters())
            flat_grad_grad_kl = torch.cat(
                [grad.contiguous().view(-1) for grad in grads])

            return flat_grad_grad_kl + v * DAMPING

        # 共轭梯度法
        step_dir = self.conjugate_gradient(
            fisher_vector_product, flat_grad.data, nsteps=10)

        # 计算自然梯度
        shs = 0.5 * (step_dir * fisher_vector_product(step_dir)
                     ).sum(0, keepdim=True)
        lm = torch.sqrt(shs / MAX_KL)
        full_step = step_dir / lm[0]

        # 更新策略网络参数
        old_params = self.get_flat_params()
        self.set_flat_params(old_params + full_step)

        # 检查KL散度约束
        if get_kl() > MAX_KL * 1.5:
            self.set_flat_params(old_params)

    def conjugate_gradient(self, f_Ax, b, nsteps, residual_tol=1e-10):
        x = torch.zeros(b.size())
        r = b.clone()
        p = b.clone()
        rdotr = torch.dot(r, r)

        for i in range(nsteps):
            Ap = f_Ax(p)
            alpha = rdotr / torch.dot(p, Ap)
            x += alpha * p
            r -= alpha * Ap
            new_rdotr = torch.dot(r, r)
            if new_rdotr < residual_tol:
                break
            p = r + (new_rdotr / rdotr) * p
            rdotr = new_rdotr

        return x

    def get_flat_params(self):
        params = []
        for param in self.policy_net.parameters():
            params.append(param.data.view(-1))
        return torch.cat(params)

    def set_flat_params(self, flat_params):
        offset = 0
        for param in self.policy_net.parameters():
            numel = param.numel()
            param.data.copy_(
                flat_params[offset:offset+numel].view(param.size()))
            offset += numel


trpo = TRPO()

print('\nTraining with TRPO...')
time.sleep(2)

for i_episode in range(400):
    states, actions, log_probs, rewards, dones, values = [], [], [], [], [], []

    s, info = env.reset()
    ep_r = 0

    while True:
        env.render()
        a, log_prob = trpo.choose_action(s)  # 进行一次行动

        # take action
        s_, r, terminated, truncated, info = env.step(a)
        done = terminated or truncated

        # modify the reward (same as before)
        x, x_dot, theta, theta_dot = s_
        r1 = (env.x_threshold - abs(x)) / env.x_threshold - 0.8
        r2 = (env.theta_threshold_radians - abs(theta)) / \
            env.theta_threshold_radians - 0.5
        r = r1 + r2

        # store transition
        states.append(s)
        actions.append(a)
        log_probs.append(log_prob.item())
        rewards.append(r)
        dones.append(done)

        # value estimate
        with torch.no_grad():
            state_tensor = torch.FloatTensor(s).unsqueeze(0)
            value = trpo.value_net(state_tensor).item()
            values.append(value)

        ep_r += r

        if done:
            # 计算最后一个状态的value
            with torch.no_grad():
                if terminated:
                    next_value = 0.0
                else:
                    next_value = trpo.value_net(
                        torch.FloatTensor(s_).unsqueeze(0)).item()

            # 计算优势函数
            advantages = trpo.compute_advantages(
                rewards, values, dones, next_value)

            # 计算回报
            returns = advantages + values

            # 更新价值网络
            trpo.update_value_net(states, returns)

            # 更新策略网络
            trpo.update_policy(states, actions, log_probs, advantages)

            print('Ep: ', i_episode, '| Ep_r: ', round(ep_r, 2))
            break

        s = s_

env.close()
View Code

 

 

posted @ 2026-03-05 23:44  阿基米德的澡盆  阅读(37)  评论(0)    收藏  举报