82 lines
3.6 KiB
Python
82 lines
3.6 KiB
Python
import torch
|
|
import torch.optim as optim
|
|
import torch.nn.functional as F
|
|
import torch.distributions as distributions
|
|
from networks import Actor, VCritic
|
|
|
|
class OffPACAgent:
|
|
def __init__(self, state_dim, action_dim, device, actor_lr=0.001, critic_lr=0.002, gamma=0.99, epsilon=0.1):
|
|
self.device = device
|
|
self.gamma = gamma
|
|
self.action_dim = action_dim
|
|
# epsilon 用于构建行为策略 beta 的探索率
|
|
self.epsilon = epsilon
|
|
|
|
# 沿用 A2C 的网络结构
|
|
self.actor = Actor(state_dim, action_dim).to(self.device)
|
|
self.critic = VCritic(state_dim).to(self.device)
|
|
|
|
self.actor_optimizer = optim.Adam(self.actor.parameters(), lr=actor_lr)
|
|
self.critic_optimizer = optim.Adam(self.critic.parameters(), lr=critic_lr)
|
|
|
|
def select_action(self, state):
|
|
# 推理时禁用梯度图计算
|
|
with torch.no_grad():
|
|
state_tensor = torch.FloatTensor(state).unsqueeze(0).to(self.device)
|
|
# 获取目标策略 pi 的动作概率分布
|
|
pi_probs = self.actor(state_tensor).squeeze(0)
|
|
|
|
# 构建行为策略 beta:在 pi 的基础上加入 epsilon 的均匀随机噪声
|
|
beta_probs = (1 - self.epsilon) * pi_probs + self.epsilon / self.action_dim
|
|
|
|
# 根据行为策略 beta 采样动作,用于与环境交互
|
|
action = distributions.Categorical(beta_probs).sample().item()
|
|
return action
|
|
|
|
def update(self, state, action, reward, next_state, next_action, done):
|
|
state_tensor = torch.FloatTensor(state).unsqueeze(0).to(self.device)
|
|
next_state_tensor = torch.FloatTensor(next_state).unsqueeze(0).to(self.device)
|
|
reward_tensor = torch.FloatTensor([reward]).unsqueeze(0).to(self.device)
|
|
action_tensor = torch.tensor([action]).to(self.device)
|
|
|
|
# --- 计算重要性采样权重 (Importance Weight) ---
|
|
# 重新获取当前最新目标策略 pi 下的动作概率
|
|
pi_probs = self.actor(state_tensor).squeeze(0)
|
|
pi_prob_a = pi_probs[action]
|
|
|
|
# 重建行为策略 beta 选出该动作的概率 (作为分母)
|
|
beta_prob_a = (1 - self.epsilon) * pi_prob_a.detach() + self.epsilon / self.action_dim
|
|
|
|
# 计算重要性权重 rho = pi(a|s) / beta(a|s)
|
|
rho = (pi_prob_a / beta_prob_a).detach()
|
|
|
|
# 【工程防崩技巧】:如果两个策略偏差过大,rho 会极大导致梯度爆炸。
|
|
# 工业界通用的做法是对 rho 进行截断 (Clip),这也是 PPO 算法的前身思想。
|
|
rho = torch.clamp(rho, 0.1, 10.0)
|
|
|
|
# --- Critic 更新 (Off-Policy) ---
|
|
v_value = self.critic(state_tensor)
|
|
next_v_value = self.critic(next_state_tensor).detach()
|
|
|
|
# 计算 TD 目标和优势函数
|
|
td_target = reward_tensor + self.gamma * next_v_value * (1 - int(done))
|
|
advantage = td_target - v_value
|
|
|
|
# Critic 损失加入了重要性权重 rho
|
|
critic_loss = (rho * F.mse_loss(v_value, td_target, reduction='none')).mean()
|
|
|
|
self.critic_optimizer.zero_grad()
|
|
critic_loss.backward()
|
|
self.critic_optimizer.step()
|
|
|
|
# --- Actor 更新 (Off-Policy) ---
|
|
action_probs = self.actor(state_tensor)
|
|
dist = distributions.Categorical(action_probs)
|
|
log_prob = dist.log_prob(action_tensor)
|
|
|
|
# Actor 梯度上升目标:rho * ln(pi) * Advantage
|
|
actor_loss = -(rho * log_prob * advantage.detach()).mean()
|
|
|
|
self.actor_optimizer.zero_grad()
|
|
self.actor_optimizer.step()
|