基于CARLA的深度强化学习TD3算法实战:从原理到自动驾驶仿真优化

1次阅读
没有评论

共计 3924 个字符,预计需要花费 10 分钟才能阅读完成。

image.webp

1. 背景与问题分析

自动驾驶决策本质上是一个连续控制问题:车辆需要根据环境状态(如周围车辆、道路标志、行人等)实时输出转向、油门和刹车等连续动作。传统强化学习算法在这一领域面临几个核心挑战:

基于 CARLA 的深度强化学习 TD3 算法实战:从原理到自动驾驶仿真优化

  • 稀疏奖励问题:安全驾驶的奖励信号往往间隔很远(如避免碰撞),导致学习效率低下
  • 高维状态空间:传感器数据(摄像头、激光雷达)维度极高,传统方法难以有效处理
  • 探索效率低:随机探索在高维连续动作空间中效果不佳,容易陷入局部最优

2. 算法选型:为什么是 TD3?

在连续控制任务中,主流算法有 DDPG、TD3 和 SAC,它们的核心区别如下:

  1. DDPG
  2. 基础确定性策略梯度算法
  3. 容易高估 Q 值导致训练不稳定
  4. 对超参数敏感

  5. TD3

  6. 双 Critic 网络缓解 Q 值高估
  7. 延迟策略更新提高稳定性
  8. 目标策略平滑减少方差

  9. SAC

  10. 随机策略更适合探索
  11. 自动调节温度系数
  12. 计算开销较大

对于自动驾驶场景,TD3 在稳定性和样本效率之间取得了较好平衡,其三大创新点特别适合本任务:

  • 双 Critic 网络:取两个 Q 网络的最小值作为更新目标,避免单一网络的高估偏差
  • 延迟策略更新:Critic 网络更新多次后才更新 Actor,保证价值估计更准确
  • 目标策略平滑:给目标动作添加噪声,防止策略陷入局部最优

3. CARLA 环境接口设计

在 CARLA 中,我们通过 PythonAPI 封装了一个符合 OpenAI Gym 接口的环境:

class CarlaEnv(gym.Env):
    def __init__(self, town='Town07'):
        self.client = carla.Client('localhost', 2000)
        self.world = self.client.load_world(town)

        # 传感器配置
        self.camera = CameraSensor(self.world)  # RGB 图像
        self.lidar = LidarSensor(self.world)    # 点云数据
        self.vehicle = VehicleController()      # 动作执行

    def step(self, action):
        # 执行动作并获取新状态
        self.vehicle.apply_control(action)
        next_state = self._get_obs()
        reward = self._calculate_reward()
        done = self._check_termination()
        return next_state, reward, done, {}

    def _get_obs(self):
        # 多模态观测拼接
        return {'camera': self.camera.get_data(),
            'lidar': self.lidar.get_data(),
            'speed': self.vehicle.get_speed()}

关键设计考虑:

  • 观测空间包含视觉(相机)、几何(LiDAR)和运动状态(速度)多模态信息
  • 奖励函数设计考虑安全性、舒适性和效率三个维度
  • 使用异步数据采集避免仿真阻塞

4. TD3 核心实现

4.1 网络架构

采用 Actor-Critic 框架,具体结构如下:

class Actor(nn.Module):
    def __init__(self, state_dim, action_dim):
        super().__init__()
        self.net = nn.Sequential(nn.Linear(state_dim, 256),
            nn.ReLU(),
            nn.Linear(256, 256),
            nn.ReLU(),
            nn.Linear(256, action_dim),
            nn.Tanh()  # 输出归一化到[-1,1]
        )

    def forward(self, state):
        return self.net(state)

class Critic(nn.Module):
    def __init__(self, state_dim, action_dim):
        super().__init__()
        # Q1 网络
        self.q1 = nn.Sequential(nn.Linear(state_dim + action_dim, 256),
            nn.ReLU(),
            nn.Linear(256, 256),
            nn.ReLU(),
            nn.Linear(256, 1)
        )
        # Q2 网络
        self.q2 = nn.Sequential(nn.Linear(state_dim + action_dim, 256),
            nn.ReLU(),
            nn.Linear(256, 256),
            nn.ReLU(),
            nn.Linear(256, 1)
        )

    def forward(self, state, action):
        sa = torch.cat([state, action], -1)
        return self.q1(sa), self.q2(sa)

4.2 关键训练逻辑

TD3 的核心训练流程如下:

  1. 采样过渡元组 (s,a,r,s’,d) 存入经验回放池
  2. 从回放池采样小批量数据
  3. 计算目标 Q 值(带平滑噪声)
  4. 更新双 Critic 网络
  5. 延迟更新 Actor 和目标网络

具体实现代码:

class TD3:
    def update(self, replay_buffer, batch_size=256):
        # 采样批次数据
        state, action, reward, next_state, done = replay_buffer.sample(batch_size)

        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) + noise).clamp(-1, 1)

            # 双 Q 目标
            target_Q1, target_Q2 = self.critic_target(next_state, next_action)
            target_Q = torch.min(target_Q1, target_Q2)
            target_Q = reward + (1 - done) * self.gamma * target_Q

        # 更新 Critic
        current_Q1, current_Q2 = self.critic(state, action)
        critic_loss = F.mse_loss(current_Q1, target_Q) + F.mse_loss(current_Q2, target_Q)
        self.critic_optimizer.zero_grad()
        critic_loss.backward()
        self.critic_optimizer.step()

        # 延迟策略更新
        if self.total_it % self.policy_freq == 0:
            actor_loss = -self.critic.Q1(state, self.actor(state)).mean()
            self.actor_optimizer.zero_grad()
            actor_loss.backward()
            self.actor_optimizer.step()

            # 软更新目标网络
            soft_update(self.critic_target, self.critic, self.tau)
            soft_update(self.actor_target, self.actor, self.tau)

5. 性能优化技巧

5.1 GPU 内存管理

CARLA 和 PyTorch 同时使用 GPU 时容易显存溢出,推荐做法:

  • 使用 torch.cuda.empty_cache() 定期清理缓存
  • 将图像数据保持在 CPU 直到需要计算
  • 使用混合精度训练

5.2 分布式采样

通过多进程并行运行多个 CARLA 实例加速数据收集:

from multiprocessing import Process, Queue

def worker(env_id, queue):
    env = CarlaEnv()
    while True:
        state = env.reset()
        done = False
        while not done:
            action = policy(state)
            next_state, reward, done, _ = env.step(action)
            queue.put((state, action, reward, next_state, done))
            state = next_state

# 启动多个工作进程
for i in range(4):
    Process(target=worker, args=(i, replay_queue)).start()

6. 常见问题解决

6.1 观测延迟问题

CARLA 的传感器数据存在异步延迟,解决方案:

  • 使用 wait_for_tick() 同步模式
  • 在状态中加入时间戳信息
  • 使用 LSTM 网络处理时序

6.2 仿真 - 现实差异

减小 domain gap 的方法:

  • 在多个 CARLA 城镇训练
  • 添加随机扰动(天气、光照等)
  • 使用领域随机化技术

7. 扩展思考

将 TD3 扩展到多车交互场景需要考虑:

  1. 集中式训练分布式执行(CTDE)架构
  2. 使用注意力机制处理可变数量邻居
  3. 设计考虑社交合规性的奖励函数

完整代码和训练结果可参考:[Colab 笔记本链接]

推荐进一步阅读:
–《Mastering Multi-Agent Reinforcement Learning》
– CARLA 官方文档中的多智能体 API
– TD3 原始论文《Addressing Function Approximation Error in Actor-Critic Methods》

通过本实践可以看到,TD3 在自动驾驶决策任务中展现出良好的稳定性和样本效率。后续可结合模仿学习进行策略初始化,或尝试将 TD3 与其他方法如 Hierarchical RL 结合处理更复杂的驾驶场景。

正文完
 0
评论(没有评论)