Files

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()