多CPU核心训练SAC智能体:OpenAI Gymnasium Bipedalwalker实现疑问
从零实现SAC+SyncVectorEnv训练BipedalWalker指南
一、先搞懂SyncVectorEnv的多进程机制
SyncVectorEnv是同步向量环境,底层会自动为每个子环境启动独立进程并调度CPU核心,不需要你手动给单个观测分配处理器——所有子环境在各自进程中并行运行step,之后同步汇总批量的观测、奖励、终止信号等结果返回给主线程。
二、核心步骤:SAC与向量环境的整合
SAC是off-policy算法,天然适配批量数据训练,核心是把向量环境返回的批量数据正确接入经验回放池和网络更新流程。
1. 向量环境初始化
通过传入环境创建函数列表来生成多并行环境,每个函数对应一个独立子环境:
import gymnasium as gym from gym.vector import SyncVectorEnv def make_env(): def _init(): env = gym.make("BipedalWalker-v3") # 可选:添加观测/奖励归一化预处理 env = gym.wrappers.NormalizeObservation(env) env = gym.wrappers.NormalizeReward(env) return env return _init # 并行环境数量根据CPU核心数调整,比如4核就设为4 num_envs = 4 vec_env = SyncVectorEnv([make_env() for _ in range(num_envs)])
2. 适配向量环境的经验回放池
向量环境返回的数据是批量维度在前的数组(如观测形状为(num_envs, obs_dim)),回放池需要支持存储批量数据并能展开为单个样本:
import numpy as np class ReplayBuffer: def __init__(self, capacity, obs_dim, act_dim): self.capacity = capacity self.obs_buf = np.zeros((capacity, obs_dim), dtype=np.float32) self.next_obs_buf = np.zeros((capacity, obs_dim), dtype=np.float32) self.acts_buf = np.zeros((capacity, act_dim), dtype=np.float32) self.rews_buf = np.zeros((capacity,), dtype=np.float32) self.dones_buf = np.zeros((capacity,), dtype=np.float32) self.ptr = 0 self.size = 0 def store_batch(self, obs, acts, rews, next_obs, dones): # 将批量数据展开为单个样本存入缓冲区 batch_size = obs.shape[0] end = self.ptr + batch_size if end > self.capacity: # 处理缓冲区溢出 remaining = self.capacity - self.ptr self.obs_buf[self.ptr:] = obs[:remaining] self.next_obs_buf[self.ptr:] = next_obs[:remaining] self.acts_buf[self.ptr:] = acts[:remaining] self.rews_buf[self.ptr:] = rews[:remaining] self.dones_buf[self.ptr:] = dones[:remaining] self.obs_buf[:batch_size-remaining] = obs[remaining:] self.next_obs_buf[:batch_size-remaining] = next_obs[remaining:] self.acts_buf[:batch_size-remaining] = acts[remaining:] self.rews_buf[:batch_size-remaining] = rews[remaining:] self.dones_buf[:batch_size-remaining] = dones[remaining:] self.ptr = batch_size - remaining else: self.obs_buf[self.ptr:end] = obs self.next_obs_buf[self.ptr:end] = next_obs self.acts_buf[self.ptr:end] = acts self.rews_buf[self.ptr:end] = rews self.dones_buf[self.ptr:end] = dones self.ptr = end self.size = min(self.size + batch_size, self.capacity) def sample(self, batch_size): idxs = np.random.randint(0, self.size, size=batch_size) return ( self.obs_buf[idxs], self.acts_buf[idxs], self.rews_buf[idxs], self.next_obs_buf[idxs], self.dones_buf[idxs] )
3. 整合SAC训练循环
核心逻辑:向量环境每一步生成批量数据存入回放池,当池内数据足够时,采样批量样本更新Actor/Critic网络:
import torch import torch.nn as nn from torch.distributions import Normal # 简化版SAC网络(可根据需求调整层数/激活函数) class Actor(nn.Module): def __init__(self, obs_dim, act_dim, act_limit): super().__init__() self.net = nn.Sequential( nn.Linear(obs_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU() ) self.mu_layer = nn.Linear(256, act_dim) self.log_std_layer = nn.Linear(256, act_dim) self.act_limit = act_limit def forward(self, obs): x = self.net(obs) mu = self.mu_layer(x) log_std = self.log_std_layer(x) log_std = torch.clamp(log_std, -20, 2) std = torch.exp(log_std) dist = Normal(mu, std) sample = dist.rsample() act = torch.tanh(sample) * self.act_limit # 修正tanh带来的概率偏差 log_prob = dist.log_prob(sample) - torch.log(1 - act.pow(2) + 1e-6) log_prob = log_prob.sum(1, keepdim=True) return act, log_prob class Critic(nn.Module): def __init__(self, obs_dim, act_dim): super().__init__() self.q1_net = nn.Sequential( nn.Linear(obs_dim + act_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 1) ) self.q2_net = nn.Sequential( nn.Linear(obs_dim + act_dim, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 1) ) def forward(self, obs, act): x = torch.cat([obs, act], dim=1) return self.q1_net(x), self.q2_net(x) # 初始化训练组件 obs_dim = vec_env.single_observation_space.shape[0] act_dim = vec_env.single_action_space.shape[0] act_limit = vec_env.single_action_space.high[0] actor = Actor(obs_dim, act_dim, act_limit) critic = Critic(obs_dim, act_dim) target_critic = Critic(obs_dim, act_dim) target_critic.load_state_dict(critic.state_dict()) actor_opt = torch.optim.Adam(actor.parameters(), lr=3e-4) critic_opt = torch.optim.Adam(critic.parameters(), lr=3e-4) replay_buffer = ReplayBuffer(capacity=1_000_000, obs_dim=obs_dim, act_dim=act_dim) gamma = 0.99 tau = 0.005 alpha = 0.2 # 熵系数,可改为自适应调整模式 # 启动训练循环 total_steps = 1_000_000 batch_size = 256 start_steps = 10_000 # 前N步随机探索 obs, _ = vec_env.reset() for step in range(total_steps): # 随机探索阶段 if step < start_steps: acts = vec_env.action_space.sample() else: obs_tensor = torch.tensor(obs, dtype=torch.float32) acts, _ = actor(obs_tensor) acts = acts.detach().numpy() # 向量环境执行step,获取批量结果 next_obs, rews, dones, truncts, _ = vec_env.step(acts) terminated = np.logical_or(dones, truncts) # 合并终止与截断信号 # 存入回放池 replay_buffer.store_batch(obs, acts, rews, next_obs, terminated) obs = next_obs # 当回放池数据足够时更新网络 if step >= start_steps: batch_obs, batch_acts, batch_rews, batch_next_obs, batch_dones = replay_buffer.sample(batch_size) # 转为PyTorch张量 batch_obs = torch.tensor(batch_obs, dtype=torch.float32) batch_acts = torch.tensor(batch_acts, dtype=torch.float32) batch_rews = torch.tensor(batch_rews, dtype=torch.float32).unsqueeze(1) batch_next_obs = torch.tensor(batch_next_obs, dtype=torch.float32) batch_dones = torch.tensor(batch_dones, dtype=torch.float32).unsqueeze(1) # 更新Critic网络 with torch.no_grad(): next_acts, next_log_probs = actor(batch_next_obs) target_q1, target_q2 = target_critic(batch_next_obs, next_acts) target_q = torch.min(target_q1, target_q2) - alpha * next_log_probs target_q = batch_rews + gamma * (1 - batch_dones) * target_q current_q1, current_q2 = critic(batch_obs, batch_acts) critic_loss = nn.MSELoss()(current_q1, target_q) + nn.MSELoss()(current_q2, target_q) critic_opt.zero_grad() critic_loss.backward() critic_opt.step() # 更新Actor网络 acts_pred, log_probs = actor(batch_obs) q1_pred, q2_pred = critic(batch_obs, acts_pred) q_pred = torch.min(q1_pred, q2_pred) actor_loss = (alpha * log_probs - q_pred).mean() actor_opt.zero_grad() actor_loss.backward() actor_opt.step() # 软更新目标Critic for t_param, param in zip(target_critic.parameters(), critic.parameters()): t_param.data.copy_(tau * param.data + (1 - tau) * t_param.data) # 定期打印训练日志 if step % 1000 == 0: print(f"Step {step}, Avg Step Reward: {np.mean(rews):.2f}")
三、关键注意事项
- 进程隔离:SyncVectorEnv的子环境在独立进程中运行,环境创建函数内不要共享全局变量,确保每个子环境独立初始化。
- 维度对齐:时刻区分向量环境(批量维度在前)和单环境的数据形状,避免维度不匹配错误。
- 熵系数优化:示例用固定alpha,实际可改为自适应熵系数(设定目标熵值为
-act_dim),提升训练稳定性。 - 预处理:BipedalWalker观测范围大,必须添加观测/奖励归一化,否则训练很难收敛。
内容的提问来源于stack exchange,提问作者Anas Rzq
相关产品推荐
相关产品推荐

