news 2026/8/8 5:38:52

强化学习QAC求最优策略的代码实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
强化学习QAC求最优策略的代码实现

理论基础:

注意:

1. get_policy_by_state_value_net() 是额外写的一个基于Q的贪心策略,不属于QAC算法,得到的策略不一定是最优的,与get_policy_by_policy_net表现一致是偶然现象。

2. 图片中的伪代码并没有说要生成多条episode,但这个为了保证每个(s,a)pair都能被访问到,会生成多条episode。

代码可运行:

import numpy as np import torch from torch import nn from env import GridWorldEnv from utils import drow_policy class QAC(object): def __init__(self, env: GridWorldEnv, gamma=0.9, lr_actor=1e-2, lr_critic=1e-2): self.env = env self.action_space_size = self.env.num_actions self.state_space_size = self.env.num_states self.gamma = gamma self.pnet = nn.Sequential( # policy_net nn.Linear(2, 16), # s -> Π(a|s) nn.ReLU(), nn.Linear(16, self.action_space_size) ) self.qnet = nn.Sequential( # q_value_net nn.Linear(2, 16), # s -> q[s,a] nn.ReLU(), nn.Linear(16, self.action_space_size) ) self.value_optimizer = torch.optim.Adam(self.qnet.parameters(), lr=lr_critic) self.policy_optimizer = torch.optim.Adam(self.pnet.parameters(), lr=lr_actor) self.policy = np.zeros((self.state_space_size, self.action_space_size)) self.q_value = np.zeros((self.state_space_size, self.action_space_size)) def decode_state(self, state): ''' :param state: int :return: 归一化后的元组 ''' i = state // self.env.size j = state % self.env.size return torch.tensor((i / (self.env.size - 1), j / (self.env.size - 1)), dtype=torch.float32) def generate_action(self, state): ''' :param state: tuple :return: int,float ''' logits = self.pnet(state) action_probs = torch.softmax(logits, dim=0) # π(a|s,θ) action_dist = torch.distributions.Categorical(action_probs) # 按分布采样 action = action_dist.sample() log_prob = action_dist.log_prob(action) # In π(a|s,θ) 注意传入的是索引,会自动做log(action_probs[action_index]) return action.item(), log_prob def solve(self, num_episodes=200): for _ in range(num_episodes): state_int = self.env.reset() state = self.decode_state(state_int) done = False while not done: action, log_prob = self.generate_action(state) # a_t,s_t,In π(a_t|s_t,θ) next_state_int, reward, done = self.env.step(state_int, action) # s_t+1,r_t+1 next_state = self.decode_state(next_state_int) if not done: next_action, _ = self.generate_action(next_state) # a_t+1 else: next_action, action_prob = None, None # Critic (value update) qvalue = self.qnet(state)[action] # q(s_t,a_t) if done: td_target = torch.tensor(reward, dtype=torch.float32) else: with torch.no_grad(): # semi gradient qvalue_next = self.qnet(next_state)[next_action] # q(s_t+1,a_t+1) td_target = torch.tensor(reward, dtype=torch.float32) + self.gamma * qvalue_next delta = td_target - qvalue # TD error self.value_optimizer.zero_grad() critic_loss = 0.5 * delta.pow(2) critic_loss.backward() self.value_optimizer.step() # Actor (policy update) qvalue = qvalue.detach() # 避免梯度污染 self.policy_optimizer.zero_grad() actor_loss = -log_prob * qvalue actor_loss.backward() self.policy_optimizer.step() state_int = next_state_int state = next_state def get_policy_by_policy_net(self): for s in range(self.state_space_size): if s in self.env.terminal: self.policy[s,4]=1 break s_t = self.decode_state(s) logits = self.pnet(s_t) action_probs = torch.softmax(logits, dim=0) a=torch.argmax(action_probs) self.policy[s,a]=1 return self.policy def get_policy_by_state_value_net(self): for s in range(self.state_space_size): if s in self.env.terminal: self.policy[s,4]=1 break a = np.argmax(self.q_value[s]) self.policy[s, a] = 1 return self.policy def get_qvalues(self): for s in range(self.state_space_size): s_t = self.decode_state(s) logits = self.qnet(s_t).detach().numpy() # q(s,a)表示在状态s执行动作a后,未来所有折扣回报的期望值,不要取softmax然后取最大 self.q_value[s, :] = logits return self.q_value if __name__ == '__main__': env = GridWorldEnv( size=5, forbidden=[(1, 2), (3, 3)], terminal=[(4, 4)], r_boundary=-1, r_other=-0.04, r_terminal=1, r_forbidden=-1, r_stay=-0.1 ) # 注意samples要大一点,否则每个state被访问到的概率很小 vi = QAC(env=env) vi.solve(num_episodes=200) print("\n state value: ") print(vi.get_qvalues()) print("\n get policy by policy net:") drow_policy(vi.get_policy_by_policy_net(), env) print("\n get policy by state value net:") drow_policy(vi.get_policy_by_state_value_net(), env)

运行结果:

1. 表现一致(终点状态不是 . 是因为没有特殊处理,其他代码保持不变。由于表现一致的情况很少,因此不再继续展示特殊处理后的输出)

2. 表现不一致

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/7 21:26:54

ZooKeeper:enableACL和requireClientSASLAuth

目录标题 🧠 一、ZooKeeper 的两个安全维度🎯 二、访问控制(ACL)1)什么是 ACL?2)ACL 相关的 Scheme(核心)3)是否开启 ACL 🔐 三、客户端认证&…

作者头像 李华
网站建设 2026/8/8 3:30:19

为什么K8s 1.24 的容器时间调整会影响宿主机的时间啊?

目录标题一、核心真相(先给结论)✅ Linux 中:二、为什么容器“有时能改时间,有时不能”?🔑 决定因素不是 K8s,而是 Linux capability三、那为什么在 K8s 1.24 更容易出现?四、K8s 1.…

作者头像 李华
网站建设 2026/8/8 4:43:13

AI时代核心竞争力:手写多智能体系统,不依赖LangChain/LlamaIndex

本文详解如何不依赖高级编排框架,使用原生Python和LLM API构建Deep Research Agent多智能体系统。系统采用反思式搜索循环和并行处理机制,实现自主规划、多轮搜索优化和结构化报告生成。文章提供完整技术实现细节、架构设计和开源代码,强调理…

作者头像 李华
网站建设 2026/8/7 11:25:35

WebSocket 对比 MQTT通信优势

——以充电桩系统为例在物联网项目中,通信协议的选择直接影响着系统的稳定性、实时性和开发效率。本文将以一个典型的充电桩系统(包含充电桩、云端服务器、微信小程序三个节点)为例,深入探讨 MQTT 和 WebSocket 两大协议的应用场景…

作者头像 李华
网站建设 2026/8/7 21:27:27

基于springboot面料花型试衣系统

基于Spring Boot的面料花型试衣系统是一个结合了后端技术和前端界面设计的综合性平台,它利用Spring Boot框架的高效性和稳定性,为用户提供了一个便捷、实时的试衣体验。以下是对该系统的详细介绍: 一、系统概述 面料花型试衣系统是一个专为面…

作者头像 李华
网站建设 2026/8/7 23:47:41

域名被污染是什么意思?还能不能继续使用?

在日常域名管理和使用过程中,不少人会遇到“域名被污染”的情况。那么,域名被污染到底是什么意思?还能否继续使用呢?一、什么是域名被污染域名被污染,通常指的是域名的解析或访问受到干扰,导致用户无法正常…

作者头像 李华