[学习笔记]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()

浙公网安备 33010602011771号