ppo和奖惩机制结合源码
用于离散动作的PPOimport torchfrom torch import nnfrom torch.nn import functional as Fimport randomimport numpy as npimport pydisimport gymfrom gym import spacesfrom torch.distributions import Categorical-------------------------------------策略网络–输出连续动作的高斯分布的均值和标准差-------------------------------------def set_seed(seed):random.seed(seed)np.random.seed(seed)torch.manual_seed(seed)if torch.cuda.is_available():torch.cuda.manual_seed(seed)torch.cuda.manual_seed_all(seed) # if you are using multi-GPU.# Ensure reproducibility in cudnntorch.backends.cudnn.deterministic Truetorch.backends.cudnn.benchmark Falseclass PolicyNet(nn.Module):definit(self, n_states, n_actions):super(PolicyNet, self).init()self.fc1 nn.Linear(n_states, 512) self.fc2 nn.Linear(512, 256) self.fc3 nn.Linear(256, 128) self.gru nn.GRU(input_size128, hidden_size128, batch_firstTrue) self.fc4 nn.Linear(128, n_actions) self._init_weights() def _init_weights(self): # 使用 Kaiming 均匀初始化方法对 全连接层 的权重进行初始化并将偏置初始化为零 nn.init.kaiming_uniform_(self.fc1.weight, nonlinearityrelu) nn.init.zeros_(self.fc1.bias) nn.init.kaiming_uniform_(self.fc2.weight, nonlinearityrelu) nn.init.zeros_(self.fc2.bias) nn.init.kaiming_uniform_(self.fc3.weight, nonlinearityrelu) nn.init.zeros_(self.fc3.bias) nn.init.kaiming_uniform_(self.fc4.weight, nonlinearityrelu) nn.init.zeros_(self.fc4.bias) nn.init.kaiming_uniform_(self.gru.weight_ih_l0, nonlinearityrelu)# 使用 Kaiming 均匀初始化方法对 GRU 层的输入权重weight_ih_l0进行初始化。 nn.init.kaiming_uniform_(self.gru.weight_hh_l0, nonlinearityrelu)# 使用 Kaiming 均匀初始化方法对 GRU 层的隐藏权重weight_hh_l0进行初始化 nn.init.zeros_(self.gru.bias_ih_l0)# 将 GRU 层的输入偏置bias_ih_l0初始化为零。 nn.init.zeros_(self.gru.bias_hh_l0)# 将 GRU 层的隐藏偏置bias_hh_l0初始化为零。 # 前向传播 def forward(self, x): # # [b, n_states] -- [b, n_hiddens] x1 F.relu(self.fc1(x)) x2 F.relu(self.fc2(x1)) x3 F.relu(self.fc3(x2)) x3 x3.unsqueeze(0) # 添加时间步维度变为 [batch_size, 1, 128] gru_out, _ self.gru(x3) # GRU的输出形状是 [batch_size, 1, 128] gru_out gru_out.squeeze(1) # 移除时间步维度变回 [batch_size, 128] x4 self.fc4(gru_out) x F.softmax(x4, dim-1) return x-------------------------------------价值网络 – 评估当前状态的价值-------------------------------------class ValueNet(nn.Module):definit(self, n_states):super(ValueNet, self).init()self.fc1 nn.Linear(n_states, 512)self.fc2 nn.Linear(512, 256)self.fc3 nn.Linear(256, 128)self.gru nn.GRU(input_size128, hidden_size128, batch_firstTrue)self.fc4 nn.Linear(128, 1)self._init_weights()def _init_weights(self): # 使用 Kaiming 均匀初始化方法对 全连接层 的权重进行初始化并将偏置初始化为零 nn.init.kaiming_uniform_(self.fc1.weight, nonlinearityrelu) nn.init.zeros_(self.fc1.bias) nn.init.kaiming_uniform_(self.fc2.weight, nonlinearityrelu) nn.init.zeros_(self.fc2.bias) nn.init.kaiming_uniform_(self.fc3.weight, nonlinearityrelu) nn.init.zeros_(self.fc3.bias) nn.init.kaiming_uniform_(self.fc4.weight, nonlinearityrelu) nn.init.zeros_(self.fc4.bias) nn.init.kaiming_uniform_(self.gru.weight_ih_l0, nonlinearityrelu)# 使用 Kaiming 均匀初始化方法对 GRU 层的输入权重weight_ih_l0进行初始化。 nn.init.kaiming_uniform_(self.gru.weight_hh_l0, nonlinearityrelu)# 使用 Kaiming 均匀初始化方法对 GRU 层的隐藏权重weight_hh_l0进行初始化 nn.init.zeros_(self.gru.bias_ih_l0)# 将 GRU 层的输入偏置bias_ih_l0初始化为零。 nn.init.zeros_(self.gru.bias_hh_l0)# 将 GRU 层的隐藏偏置bias_hh_l0初始化为零。 # 前向传播 def forward(self, x): x1 F.relu(self.fc1(x)) # [b,n_states]--[b,n_hiddens] x2 F.relu(self.fc2(x1)) x3 F.relu(self.fc3(x2)) x3 x3.unsqueeze(0) # 添加时间步维度变为 [batch_size, 1, 128] gru_out, _ self.gru(x3) # GRU的输出形状是 [batch_size, 1, 128] gru_out gru_out.squeeze(1) # 移除时间步维度变回 [batch_size, 128] x4 self.fc4(gru_out) return x4-------------------------------------模型构建–处理连续动作-------------------------------------class PPO:definit(self, n_states,n_actions,actor_lr,critic_lr,device):# 实例化策略网络self.actor PolicyNet(n_states, n_actions).to(device)self.actor_old PolicyNet(n_states, n_actions).to(device)# 实例化价值网络self.critic ValueNet(n_states).to(device)# 策略网络的优化器self.actor_optimizer torch.optim.Adam(self.actor.parameters(), lractor_lr)# 价值网络的优化器self.critic_optimizer torch.optim.Adam(self.critic.parameters(), lrcritic_lr)# 属性分配 self.lmbda 0.9 # GAE优势函数的缩放因子 self.epochs 5 # 一条序列的数据用来训练多少轮 self.eps 0.2 # 截断范围 self.gamma 0.9 # 折扣系数 self.device device # 动作选择 def take_action(self, state): # 输入当前时刻的状态 # [n_states]--[1,n_states]--tensor state torch.tensor(state,dtypetorch.float32).to(self.device) # 预测当前状态的动作输出动作概率的概率分布 probs self.actor(state) # 输出在每个动作上的概率 m Categorical(probs)# 使用Categorical分布将输出的概率转换为一个概率分布对象m action m.sample()# 通过调用m.sample()方法从该分布中采样一个动作。 # 随机选择动作 return action.item() # 返回动作值 # 训练 def update(self,states,next_states, actions, rewards): self.actor_old.load_state_dict(self.actor.state_dict()) states torch.FloatTensor(states) next_states torch.FloatTensor(next_states) actions torch.LongTensor(actions) rewards torch.FloatTensor(rewards) next_states_target self.critic(next_states) # 价值网络--目标当前时刻的state_value [b,1] td_target rewards self.gamma * next_states_target.squeeze() # 价值网络--预测当前时刻的state_value [b,n_states]--[b,1] td_value self.critic(states) # 时序差分预测值-目标值 # [b,1] td_delta td_target - td_value.squeeze() # 对时序差分结果计算GAE优势函数 td_delta td_delta.cpu().detach().numpy() # [b,1] advantage_list [] # 保存每个时刻的优势函数 advantage 0 # 优势函数初始值 # 逆序遍历时序差分结果把最后一时刻的放前面 for delta in td_delta[::-1]: advantage self.gamma * self.lmbda * advantage delta advantage_list.append(advantage) # 正序排列优势函数 advantage_list.reverse() # numpy -- tensor advantage torch.tensor(advantage_list, dtypetorch.float).to(self.device) old_probs self.actor_old(states) # 输出在每个动作上的概率 old_dist Categorical(old_probs) # 策略网络--预测当前状态选择的动作的高斯分布 # 一个序列训练epochs次 old_log_probs old_dist.log_prob(actions) for _ in range(self.epochs): new_probs self.actor(states)# 对状态集合进行预测得到每个动作的概率分布new_probs new_dist Categorical(new_probs) # 使用Categorical分布将输出的概率转换为一个概率分布对象new_dist new_log_probs new_dist.log_prob(actions)# 计算每个动作的对数概率 # action_dists torch.distributions.Normal(mu, std) # # 当前策略在 t 时刻智能体处于状态 s 所采取动作的行为概率 # log_prob action_dists.log_prob(actions)# 新策略 # # 计算概率的比值来控制新策略更新幅度 ratio torch.exp(new_log_probs - old_log_probs) # 公式的左侧项 surr1 ratio * advantage # 公式的右侧项截断 surr2 torch.clamp(ratio, 1 - self.eps, 1 self.eps) * advantage # 策略网络的损失PPO-clip actor_loss torch.mean(-torch.min(surr1, surr2)) # 价值网络的当前时刻预测值与目标价值网络当前时刻的state_value之差 critic_loss torch.mean(F.mse_loss(self.critic(states).squeeze(),td_target.detach())) # 优化器清0 self.actor_optimizer.zero_grad() self.critic_optimizer.zero_grad() # 梯度反传 actor_loss.backward(retain_graphTrue) critic_loss.backward(retain_graphTrue) # 参数更新 self.actor_optimizer.step() self.critic_optimizer.step() def save(self,episode,moudle_dir): torch.save(self.actor.state_dict(), f{moudle_dir}/{episode}PPO_actor_dec.pth) torch.save(self.critic.state_dict(), f{moudle_dir}/{episode}PPO_critic_dec.pth) print(...save model...) def load(self,moudle_dir): self.actor.load_state_dict(torch.load(f{moudle_dir}/PPO_actor_dec.pth)) self.critic.load_state_dict(torch.load(f{moudle_dir}/PPO_critic_dec.pth)) print(...load...)class AFSIMEnv(gym.Env):definit(self, entity_id):super(AFSIMEnv, self).init()self.entity_id entity_idself.dis_network pydis.Network()# 定义观测空间和动作空间 self.observation_space spaces.Box(low-1, high1, shape(5,)) self.action_space spaces.Box(low-1, high1, shape(2,)) def reset(self): # 重置环境初始化实体状态 initial_state self._get_entity_state() return np.array(initial_state, dtypenp.float32) def _get_entity_state(self): # 通过DIS获取实体状态 entity_state_pdu self.dis_network.receive_pdu() if entity_state_pdu is None: return [0, 0, 0, 0, 0] # 获取实体的位置信息和姿态信息作为观测值 pos_x, pos_y, pos_z entity_state_pdu.entity_location.x, entity_state_pdu.entity_location.y, entity_state_pdu.entity_location.z pitch, roll, yaw entity_state_pdu.entity_orientation.pitch, entity_state_pdu.entity_orientation.roll, entity_state_pdu.entity_orientation.yaw return [pos_x, pos_y, pos_z, pitch, yaw] def step(self, action): # 通过DIS向AFSIM发送控制命令 control_pdu pydis.EntityStatePDU() control_pdu.entity_id self.entity_id control_pdu.entity_orientation.pitch action[0] control_pdu.entity_orientation.yaw action[1] self.dis_network.send_pdu(control_pdu) # 获取新的状态和奖励 state self._get_entity_state() reward -np.linalg.norm(state[:2]) # 假设奖励为距离目标点的负距离 done bool(state[0] 0.1 and state[1] 0.1) return np.array(state, dtypenp.float32), reward, done, {}在这里插入代码片

相关新闻

从学术研究到商业应用:SenseNova-U1.5-8B-MoT-Preview的技术演进与产业价值

从学术研究到商业应用:SenseNova-U1.5-8B-MoT-Preview的技术演进与产业价值

从学术研究到商业应用:SenseNova-U1.5-8B-MoT-Preview的技术演进与产业价值 【免费下载链接】SenseNova-U1.5-8B-MoT-Preview 项目地址: https://ai.gitcode.com/SenseNova/SenseNova-U1.5-8B-MoT-Preview SenseNova-U1.5-8B-MoT-Preview是基于NEO-unify架构…

2026/10/5 18:50:50 阅读更多 →
一图四席!派拉软件入选《2026人工智能安全产业全景图》四大核心赛道

一图四席!派拉软件入选《2026人工智能安全产业全景图》四大核心赛道

近日,由上海市信息安全行业协会与安在新媒体联合编制的《2026中国AI安全产业发展调查报告》正式全网首发,并同步推出年度重磅成果——《2026人工智能安全产业全景图》。作为研判国内AI安全产业格局的权威参考,这份全景图覆盖AI安全全产业链核…

2026/10/2 4:58:57 阅读更多 →
中小企业怎么选智能客服?2026年saas智能客服选购指南分享

中小企业怎么选智能客服?2026年saas智能客服选购指南分享

很多中小企业管理者在引入智能客服时,都会陷入同样的困惑:市面上产品不少,宣传各有侧重,到底哪一款才适合自己? 选型过程中,有三类问题尤其常见 第一类:看参数不看匹配。 不少企业选型时先看技…

2026/10/7 12:39:03 阅读更多 →

最新新闻

YOLOv8水下管道检测识别:从数据集整理到模型训练与部署全流程解析

YOLOv8水下管道检测识别:从数据集整理到模型训练与部署全流程解析

简介:面向海洋工程与基础设施巡检场景,这套基于YOLOv8的水下管道检测资料包,为需要快速落地目标检测方案的开发者和巡检人员,提供了从数据集到训练模型再到部署参考的一站式支持。压缩包共2000个文件,其中以1985个VOC格…

2026/10/11 0:34:55 阅读更多 →
痤疮检测数据集VOC转YOLO实战:915张图像训练YOLOv8全流程

痤疮检测数据集VOC转YOLO实战:915张图像训练YOLOv8全流程

简介:面向目标检测与医学图像识别场景的痤疮检测数据集,共915张标注图像,采用Pascal VOC与YOLO双格式存储,配套xml和txt标注文件,适合希望直接训练YOLO系列、SSD、Faster R-CNN等模型的开发者或研究人员使用。类别仅含…

2026/10/11 0:34:55 阅读更多 →
Flink CDC 2.3.0实战:SQL Server实时同步到MySQL全指南

Flink CDC 2.3.0实战:SQL Server实时同步到MySQL全指南

简介:面向大数据实时同步场景,这套工程基于 flink-connector-sqlserver-cdc 2.3.0,提供从 SQL Server 到 MySQL 的完整 CDC 同步示例,适合具备 Flink 基础、正在搭建实时数仓或跨库同步链路的开发者。包体共 21 个文件&#xff0c…

2026/10/11 0:34:55 阅读更多 →
轻量级实时人眼状态识别Python实现

轻量级实时人眼状态识别Python实现

简介:本资源是一套基于Python与OpenCV实现的实时人眼识别与眨眼/闭眼状态检测的完整开发方案,面向计算机视觉初学者、人工智能课程学习者及人脸交互项目开发者,解决疲劳监测、注意力评估、人机交互等实际场景中的关键感知需求。压缩包共61个文…

2026/10/11 0:34:55 阅读更多 →
Python基础语法一站式指南:从环境搭建到项目实战

Python基础语法一站式指南:从环境搭建到项目实战

很多朋友学Python时,最大的障碍其实不是"某个语法不会",而是知识点太碎,学完列表学字典,学完函数学文件,一到自己写项目就全乱套。我这些年用Python写自动化脚本、做数据分析、折腾小爬虫,最大的体会是:基础语法必须串成一条线来学。这篇指南我打算从装环境一路写到函…

2026/10/11 0:33:54 阅读更多 →
基于MPC的混合储能微电网双层能量管理:Matlab仿真与实现

基于MPC的混合储能微电网双层能量管理:Matlab仿真与实现

1. 整体设计与思路拆解:为什么混合储能微电网需要“双层”和“MPC”先聊一个经常被问到的问题:既然已经有能量管理系统(EMS)了,为什么还要搞“双层”?又为什么要专门上模型预测控制(MPC&#xf…

2026/10/11 0:33:54 阅读更多 →

日新闻

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

简介:基于 ARIMA、LSTM、Transformer 等模型的流感时间序列预测 Python 源码,面向计算机相关专业课程设计与期末大作业学生,以及项目实战学习者。内容覆盖预处理、平稳性检验、定阶、残差分析、多模型对比预测的完整时序建模流程,…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别 做影刀RPA自动化,十个新手有八个栽在"往输入框里填东西"这件事上:要么填不进去,要么填了一半,要么直接把原来内容追加在后面。这背后的根因&…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容 1. 认识影刀:什么场景该用RPA采小说数据 起点中文网的页面结构相对稳定——分类榜单、书籍详情、章节内容三块独立页面,跳转链路清晰。这种场景非常适合影刀自动化&#x…

2026/10/11 0:00:27 阅读更多 →

周新闻

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

简介:基于 ARIMA、LSTM、Transformer 等模型的流感时间序列预测 Python 源码,面向计算机相关专业课程设计与期末大作业学生,以及项目实战学习者。内容覆盖预处理、平稳性检验、定阶、残差分析、多模型对比预测的完整时序建模流程,…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别 做影刀RPA自动化,十个新手有八个栽在"往输入框里填东西"这件事上:要么填不进去,要么填了一半,要么直接把原来内容追加在后面。这背后的根因&…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容 1. 认识影刀:什么场景该用RPA采小说数据 起点中文网的页面结构相对稳定——分类榜单、书籍详情、章节内容三块独立页面,跳转链路清晰。这种场景非常适合影刀自动化&#x…

2026/10/11 0:00:27 阅读更多 →

月新闻

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/10 5:23:50 阅读更多 →
Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/9 21:32:20 阅读更多 →
黑夜航拍船只数据集训练YOLOV5模型全流程解析

黑夜航拍船只数据集训练YOLOV5模型全流程解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/10 10:38:42 阅读更多 →