TD3-HER算法Actor Loss无法下降及智能体停滞问题排查
Carla Endless-v0场景下TD3-HER智能体训练异常
问题详情
在Carla环境的Endless-v0场景中训练基于TD3-HER的智能体时,出现以下异常:
- Actor和Critic的损失曲线异常(见下图)
- 智能体训练初期会尝试探索驾驶方式,但后续陷入停滞不再移动
- 已调整奖励函数、经验池容量及学习率,但问题仍未解决

模型代码
import torch import torch.nn as nn import torch.nn.functional as F import numpy as np from collections import deque import random from models.encoder.resnet18 import ResNet18Encoder import os import itertools import sys class Actor(nn.Module): def __init__(self, encoder, goal_dim, action_dim, vector_dim, max_action): super(Actor, self).__init__() self.encoder = encoder self.max_action = max_action self.mlp = nn.Sequential( nn.Linear(self.encoder.feature_dim + goal_dim + vector_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, action_dim), nn.Tanh() ) def forward(self, state_img, state_vector, goal): state_features = self.encoder(state_img) combined = torch.cat([state_features, state_vector, goal], dim=1) return self.max_action * self.mlp(combined) class Critic1(nn.Module): def __init__(self, encoder, goal_dim, action_dim, vector_dim): super(Critic1, self).__init__() self.encoder = encoder self.mlp = nn.Sequential( nn.Linear(self.encoder.feature_dim + goal_dim + action_dim + vector_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 1) ) def forward(self, state_img, state_vector, goal, action): state_features = self.encoder(state_img) combined = torch.cat([state_features, state_vector, goal, action], dim=1) q_value = self.mlp(combined) return q_value class Critic2(nn.Module): def __init__(self, encoder, goal_dim, action_dim, vector_dim): super(Critic2, self).__init__() self.encoder = encoder self.mlp = nn.Sequential( nn.Linear(self.encoder.feature_dim + goal_dim + action_dim + vector_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 1) ) def forward(self, state_img, state_vector, goal, action): state_features = self.encoder(state_img) combined = torch.cat([state_features, state_vector, goal, action], dim=1) q_value = self.mlp(combined) return q_value class HerReplayBuffer: def __init__(self, capacity, k_future): self.capacity = capacity self.buffer = deque(maxlen=capacity) self.future_p = 1 - (1. / (1 + k_future)) def add(self, episode_transitions): self.buffer.append(episode_transitions) def sample(self, batch_size): ep_indices = np.random.randint(0, len(self.buffer), batch_size) transitions = [] for idx in ep_indices: episode = self.buffer[idx] t_sample = np.random.randint(0, len(episode)) original_transition = episode[t_sample] if np.random.uniform() < self.future_p: future_t = np.random.randint(t_sample, len(episode)) future_ag = episode[future_t]['ag_next'] her_transition = original_transition.copy() her_transition['g'] = future_ag her_transition['r'] = self.compute_reward(her_transition['ag_next'], her_transition['g']) transitions.append(her_transition) else: transitions.append(original_transition) obs_img = np.array([t['obs']['image'] for t in transitions]) obs_vec = np.array([t['obs']['vector'] for t in transitions]) actions = np.array([t['a'] for t in transitions]) rewards = np.array([t['r'] for t in transitions]).reshape(-1, 1) next_obs_img = np.array([t['obs_next']['image'] for t in transitions]) next_obs_vec = np.array([t['obs_next']['vector'] for t in transitions]) g = np.array([t['g'] for t in transitions]) return obs_img, obs_vec, actions, rewards, next_obs_img, next_obs_vec, g @staticmethod def compute_reward(achieved_goal, desired_goal, threshold=2.5): distance = np.linalg.norm(achieved_goal - desired_goal, axis=-1) return -distance def __len__(self): return sum(len(episode) for episode in self.buffer) class TD3_HER_Agent: def __init__(self, goal_dim, action_dim, vector_dim, max_action, k_future=4, discount=0.99, tau=0.005, policy_noise=0.2, noise_clip=0.5, policy_freq=2): self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") encoder = ResNet18Encoder().to(self.device) encoder_target = ResNet18Encoder().to(self.device) encoder_target.load_state_dict(encoder.state_dict()) self.actor = Actor(encoder, goal_dim, action_dim, vector_dim, max_action).to(self.device) self.actor_target = Actor(encoder_target, goal_dim, action_dim, vector_dim, max_action).to(self.device) self.actor_target.load_state_dict(self.actor.state_dict()) self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=3e-5) self.critic1 = Critic1(encoder, goal_dim, action_dim, vector_dim).to(self.device) self.critic2 = Critic2(encoder, goal_dim, action_dim, vector_dim).to(self.device) self.critic1_target = Critic1(encoder_target, goal_dim, action_dim, vector_dim).to(self.device) self.critic2_target = Critic2(encoder_target, goal_dim, action_dim, vector_dim).to(self.device) self.critic1_target.load_state_dict(self.critic1.state_dict()) self.critic2_target.load_state_dict(self.critic2.state_dict()) self.critic_optimizer = torch.optim.Adam( itertools.chain(self.critic1.parameters(), self.critic2.parameters()), lr=3e-4 ) self.replay_buffer = HerReplayBuffer(capacity=500000, k_future=k_future) self.max_action = max_action self.discount = discount self.tau = tau self.policy_noise = policy_noise self.noise_clip = noise_clip self.policy_freq = policy_freq self.total_it = 0 def select_action(self, state_img, state_vector, goal, noise_scale=0.3): state_img = torch.FloatTensor(state_img).unsqueeze(0).to(self.device) state_vector = torch.FloatTensor(state_vector).unsqueeze(0).to(self.device) goal = torch.FloatTensor(goal).unsqueeze(0).to(self.device) self.actor.eval() with torch.no_grad(): action = self.actor(state_img, state_vector, goal).cpu().data.numpy().flatten() self.actor.train() noise = np.random.normal(0, self.max_action * noise_scale, size=action.shape) action = (action + noise).clip(-self.max_action, self.max_action) return action def train(self, batch_size=256): self.total_it += 1 if len(self.replay_buffer) < batch_size: return None, None obs_img, obs_vec, action, reward, next_obs_img, next_obs_vec, goal = self.replay_buffer.sample(batch_size) state_img = torch.FloatTensor(obs_img).to(self.device) state_vec = torch.FloatTensor(obs_vec).to(self.device) action = torch.FloatTensor(action).to(self.device) reward = torch.FloatTensor(reward).to(self.device) next_state_img = torch.FloatTensor(next_obs_img).to(self.device) next_state_vec = torch.FloatTensor(next_obs_vec).to(self.device) goal = torch.FloatTensor(goal).to(self.device) with torch.no_grad(): noise = (torch.randn_like(action) * self.policy_noise).clamp(-self.noise_clip, self.noise_clip) next_action = (self.actor_target(next_state_img, next_state_vec, goal) + noise).clamp(-self.max_action, self.max_action) target_Q1 = self.critic1_target(next_state_img, next_state_vec, goal, next_action) target_Q2 = self.critic2_target(next_state_img, next_state_vec, goal, next_action) target_Q = torch.min(target_Q1, target_Q2) target_Q = reward + self.discount * target_Q current_Q1 = self.critic1(state_img, state_vec, goal, action) current_Q2 = self.critic2(state_img, state_vec, goal, action) critic_loss = F.smooth_l1_loss(current_Q1, target_Q.detach()) + F.smooth_l1_loss(current_Q2, target_Q.detach()) self.critic_optimizer.zero_grad() critic_loss.backward() torch.nn.utils.clip_grad_norm_(self.critic1.parameters(), 10.0) torch.nn.utils.clip_grad_norm_(self.critic2.parameters(), 10.0) self.critic_optimizer.step() actor_loss = None if self.total_it % self.policy_freq == 0: actor_action = self.actor(state_img, state_vec, goal) actor_loss = -self.critic1(state_img.detach(), state_vec.detach(), goal.detach(), actor_action).mean() self.actor_optimizer.zero_grad() actor_loss.backward() torch.nn.utils.clip_grad_norm_(self.actor.parameters(), max_norm=1.0) self.actor_optimizer.step() for param, target_param in zip(self.critic1.parameters(), self.critic1_target.parameters()): target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data) for param, target_param in zip(self.critic2.parameters(), self.critic2_target.parameters()): target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data) for param, target_param in zip(self.actor.parameters(), self.actor_target.parameters()): target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data) return actor_loss.item() if actor_loss is not None else None, critic_loss.item() def save(self, directory): os.makedirs(directory, exist_ok=True) actor_path = os.path.join(directory, "actor.pt") critic1_path = os.path.join(directory, "critic1.pt") critic2_path = os.path.join(directory, "critic2.pt") torch.save(self.actor.state_dict(), actor_path) torch.save(self.critic1.state_dict(), critic1_path) torch.save(self.critic2.state_dict(), critic2_path) print(f"--- Model saved to directory: {directory} ---") def load(self, directory): actor_path = os.path.join(directory, "actor.pt") critic1_path = os.path.join(directory, "critic1.pt") critic2_path = os.path.join(directory, "critic2.pt") self.actor.load_state_dict(torch.load(actor_path, map_location=self.device)) self.critic1.load_state_dict(torch.load(critic1_path, map_location=self.device)) self.critic2.load_state_dict(torch.load(critic2_path, map_location=self.device)) self.actor_target.load_state_dict(self.actor.state_dict()) self.critic1_target.load_state_dict(self.critic1.state_dict()) self.critic2_target.load_state_dict(self.critic2.state_dict()) print(f"--- Model loaded from directory: {directory} ---")
内容的提问来源于stack exchange,提问作者Jiashu Li
相关产品推荐
相关产品推荐

