代码已经发布在https://github.com/123sdadaw/PPO_lagrangian

最近对安全强化学习(safe RL)比较感兴趣,不过发现针对于安全强化学习的代码开源较少,且各有各的实现风格,于是自己针对PPO编写了PPO lagrangian,特地分享出来与大家探讨学习。代码已发布在https://github.com/123sdadaw/PPO_lagrangian

一、ppo lagrangian的原理

PPO的原理就不赘述了,PPO lagrangian的原理可以用下面一张图解释:

在强化学习中,智能体通过与环境交互执行动作并收到奖励(reward),用于更新价值函数。安全强化学习(Safe RL)在此基础上引入了成本(cost)信号,评估智能体在执行动作时是否违反安全约束,因此与环境交互会返回两个值,reward和cost,其中cost和reward一样,由自定义的约束函数计算

标准强化学习中,奖励信号用于更新价值函数;而安全强化学习则额外更新安全价值函数,通过成本来评估安全性。在PPO-Lagrangian算法中,增加了一个安全价值网络(Safe Critic)来评估成本(cost),并将网络输出结果纳入演员网络(actor)的损失函数中,从而在优化奖励的同时,确保满足安全约束。

此外,PPO-Lagrangian使用拉格朗日乘子(Lagrange multiplier)来动态调整安全约束。当成本较高时,拉格朗日乘子增大,从而加强对演员网络更新的约束。

二、代码细节

这里仅对于PPO lagrangian与标准PPO不同之处做出说明。

注:基础PPO代码参考自https://github.com/Lizhi-sjtu/DRL-code-pytorch,知乎为https://zhuanlan.zhihu.com/p/512327050

1.网络定义

标准PPO只有actor和critic两个网络,PPO lagrangian添加的safe critic网络与critic网络结构相同,如下:

#############################################
# 定义 Critic 网络(用于奖励值估计)
#############################################
class Critic(nn.Module):
    def __init__(self, args):
        super(Critic, self).__init__()
        self.fc1 = nn.Linear(args.state_dim, args.hidden_width)
        self.fc2 = nn.Linear(args.hidden_width, args.hidden_width)
        self.fc3 = nn.Linear(args.hidden_width, 1)
        self.activate_func = [nn.ReLU(), nn.Tanh()][args.use_tanh]  # Trick10: use tanh

        if args.use_orthogonal_init:
            print("------use_orthogonal_init------")
            orthogonal_init(self.fc1)
            orthogonal_init(self.fc2)
            orthogonal_init(self.fc3)

    def forward(self, s):
        s = self.activate_func(self.fc1(s))
        s = self.activate_func(self.fc2(s))
        v_s = self.fc3(s)
        return v_s


#############################################
# 定义 SafeCritic 网络(用于安全代价估计)
#############################################
class SafeCritic(nn.Module):
    def __init__(self, args):
        super(SafeCritic, self).__init__()
        self.fc1 = nn.Linear(args.state_dim, args.hidden_width)
        self.fc2 = nn.Linear(args.hidden_width, args.hidden_width)
        self.fc3 = nn.Linear(args.hidden_width, 1)
        self.activate_func = [nn.ReLU(), nn.Tanh()][args.use_tanh] # Trick10: use tanh

        if args.use_orthogonal_init:
            print("------use_orthogonal_init (SafeCritic)------")
            orthogonal_init(self.fc1)
            orthogonal_init(self.fc2)
            orthogonal_init(self.fc3)

    def forward(self, s):
        s = self.activate_func(self.fc1(s))
        s = self.activate_func(self.fc2(s))
        cost_value = self.fc3(s)
        return cost_value

2.广义优势计算

cost和reward相同。都要使用GAE来计算优势:

        adv = []
        gae = 0
        with torch.no_grad():  # adv and v_target have no gradient
            vs = self.critic(s)
            vs_ = self.critic(s_)
            deltas = r + self.gamma * (1.0 - dw) * vs_ - vs   # 计算TD误差
            for delta, d in zip(reversed(deltas.flatten().numpy()), reversed(done.flatten().numpy())):
                gae = delta + self.gamma * self.lamda * gae * (1.0 - d)
                adv.insert(0, gae)
            adv = torch.tensor(adv, dtype=torch.float).view(-1, 1)
            v_target = adv + vs  # 计算状态值目标
            if self.use_adv_norm:  # Trick 1:advantage normalization
                adv = ((adv - adv.mean()) / (adv.std() + 1e-5))

        # 计算安全代价优势 cost_adv 和安全目标值 cost_target(使用 GAE)
        cost_adv = []
        gae_c = 0
        with torch.no_grad():
            cost_values = self.safe_critic(s)
            cost_values_next = self.safe_critic(s_)
            cost_deltas = c + self.cost_gamma * (1.0 - dw) * cost_values_next - cost_values
            for delta, d in zip(reversed(cost_deltas.flatten().cpu().numpy()), reversed(done.flatten().cpu().numpy())):
                gae_c = delta + self.cost_gamma * self.cost_lamda * gae_c * (1.0 - d)
                cost_adv.insert(0, gae_c)
            cost_adv = torch.tensor(cost_adv, dtype=torch.float).view(-1, 1)
            cost_target = cost_adv + cost_values
            if self.use_adv_norm:  # Trick 1:advantage normalization
                cost_adv = ((cost_adv - cost_adv.mean()) / (cost_adv.std() + 1e-5))

3.safe critic更新

与critic更新流程相同:

# Update critic
v_s = self.critic(s[index])
critic_loss = F.mse_loss(v_target[index], v_s)
self.optimizer_critic.zero_grad()
critic_loss.backward()
if self.use_grad_clip:  # Trick 7: Gradient clip
    torch.nn.utils.clip_grad_norm_(self.critic.parameters(), 0.5)
self.optimizer_critic.step()

# 更新安全 critic(安全代价值函数)的损失
v_cost = self.safe_critic(s[index])
safe_critic_loss = F.mse_loss(cost_target[index], v_cost)
self.optimizer_safe_critic.zero_grad()
safe_critic_loss.backward()
if self.use_grad_clip:
    torch.nn.utils.clip_grad_norm_(self.safe_critic.parameters(), 0.5)
self.optimizer_safe_critic.step()

4.actor更新

与标准PPO不同,actor更新时,损失加入了拉格朗日项,如下:

surr1 = ratios * adv[index]  # Only calculate the gradient of 'a_logprob_now' in ratios
surr2 = torch.clamp(ratios, 1 - self.epsilon, 1 + self.epsilon) * adv[index]
actor_loss = -torch.min(surr1, surr2) - self.entropy_coef * dist_entropy  # shape(mini_batch_size X 1)

# 安全部分(这里不进行 clip,直接乘以 cost_adv)
actor_loss_cost = ratios * cost_adv[index]

# 加入安全部分后的actor损失
actor_loss = actor_loss + self.lambda_cost * actor_loss_cost

注意这里拉格朗日惩罚项使用的是加号,而数学公式中写的是减号,这是因为强化学习的目标一般是最大化reward,而pytorch中的优化器一般是最小化loss,因此需要整体加入负号。

5.拉格朗日乘子的初始化及更新

拉格朗日乘子初始化为1:

# 初始化拉格朗日乘子(确保非负)
self.lambda_cost = torch.tensor(1.0, requires_grad=True)
self.optimizer_lambda = torch.optim.Adam([self.lambda_cost], lr=self.lr_multiplier)

拉格朗日乘子更新过程为(注意乘子更新时要保持非负):

# 更新拉格朗日乘子 lambda_cost
# 计算整个 batch 上的平均安全代价优势与安全阈值的偏差
cost_violation = cost_adv.mean() - self.cost_limit
lambda_loss = - self.lambda_cost * cost_violation.detach()  # 注意 detach,避免影响其他梯度
self.optimizer_lambda.zero_grad()
lambda_loss.backward()
self.optimizer_lambda.step()
# 保证 lambda_cost 非负
with torch.no_grad():
    self.lambda_cost.clamp_(0)

其中cost_limit一般设置为0,即满足约束时,cost为0,不满足约束,cost是一个大于0的值。


总结

希望能帮到大家!GitHub上的代码大家觉得好,也可以多多star!

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐