Files
RL-Study/RL_Algothrithms/networks.py
T

44 lines
1.3 KiB
Python
Raw Normal View History

2026-03-18 09:07:01 +00:00
import torch
import torch.nn as nn
import torch.nn.functional as F
# 策略网络 (离散动作 Actor)
class Actor(nn.Module):
def __init__(self, state_dim, action_dim):
super(Actor, self).__init__()
# 定义两层隐藏层
self.fc1 = nn.Linear(state_dim, 128)
self.fc2 = nn.Linear(128, action_dim)
def forward(self, state):
x = F.relu(self.fc1(state))
# 输出动作的概率分布
action_probs = F.softmax(self.fc2(x), dim=-1)
return action_probs
# Q值网络 (用于 QAC)
class QCritic(nn.Module):
def __init__(self, state_dim, action_dim):
super(QCritic, self).__init__()
self.fc1 = nn.Linear(state_dim, 128)
self.fc2 = nn.Linear(128, action_dim)
def forward(self, state):
x = F.relu(self.fc1(state))
# 输出各个动作的具体价值
q_values = self.fc2(x)
return q_values
# V值网络 (用于 A2C,由于加入了基线,只需输出状态价值标量)
class VCritic(nn.Module):
def __init__(self, state_dim):
super(VCritic, self).__init__()
self.fc1 = nn.Linear(state_dim, 128)
self.fc2 = nn.Linear(128, 1)
def forward(self, state):
x = F.relu(self.fc1(state))
# 输出当前状态的价值评估
v_value = self.fc2(x)
return v_value