Files

75 lines
2.6 KiB
Python

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributions import Normal
class ActorNet(nn.Module):
"""
策略网络 (Actor):输入环境状态,输出连续动作的高斯分布参数 (均值和标准差)
"""
def __init__(self, state_dim, action_dim, hidden_dim=64):
super(ActorNet, self).__init__()
# 定义两层隐藏层,提取状态特征
self.fc1 = nn.Linear(state_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, hidden_dim)
# 输出层:计算动作的均值 (Mean)
self.mean_layer = nn.Linear(hidden_dim, action_dim)
# 定义对数标准差 (Log Standard Deviation) 为可训练的独立参数
# 注意:在 TRPO/PPO 中,通常让标准差独立于状态,这能让训练更稳定
# 初始值设为 0,意味着初始标准差为 exp(0) = 1.0
self.log_std = nn.Parameter(torch.zeros(1, action_dim))
def forward(self, state):
# 前向传播提取特征
x = F.tanh(self.fc1(state))
x = F.tanh(self.fc2(x))
# 计算均值
mean = self.mean_layer(x)
# 将对数标准差扩展到与 batch size 相同的维度
log_std = self.log_std.expand_as(mean)
# 转化为标准差
std = torch.exp(log_std)
return mean, std
def get_action(self, state):
"""
根据当前状态采样动作,并返回对应的对数概率 (log probability)
"""
# 获取当前状态的均值和标准差
mean, std = self.forward(state)
# 构建正态分布 (高斯分布)
dist = Normal(mean, std)
# 从分布中采样一个动作
action = dist.sample()
# 返回采样动作及其对数概率 (后续计算重要性采样权重时必须用到)
return action, dist.log_prob(action)
class CriticNet(nn.Module):
"""
价值网络 (Critic):输入环境状态,评估该状态的预期收益 (标量 V 值)
"""
def __init__(self, state_dim, hidden_dim=64):
super(CriticNet, self).__init__()
# 定义隐藏层
self.fc1 = nn.Linear(state_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, hidden_dim)
# 输出层:计算状态价值 V(s),输出维度为 1
self.value_layer = nn.Linear(hidden_dim, 1)
def forward(self, state):
# 前向传播计算价值
x = F.tanh(self.fc1(state))
x = F.tanh(self.fc2(x))
value = self.value_layer(x)
return value