4dbea5f0a6c33229f44a4b4c2b7a80639206062c
MAGAIL4AutoDrive
基于多智能体生成对抗模仿学习(MAGAIL)的自动驾驶训练系统 | MetaDrive + Waymo Open Motion Dataset
本项目利用 Waymo 真实驾驶数据,通过 MetaDrive 仿真环境构建专家回放系统,提取车辆状态与动作,用于训练多智能体模仿学习算法 (MAGAIL)。
📁 核心模块
Env/expert_replay_env.py: 专家回放环境。核心类ExpertReplayEnv,负责读取 Waymo 轨迹,计算逆动力学动作,并过滤非道路/静态车辆。Env/inverse_dynamics.py: 逆动力学模块。根据车辆位置和航向计算油门、刹车和转向动作。scripts/generate_expert_data.py: 数据收集脚本。批量运行场景并保存训练数据。scripts/visualize_replay.py: 可视化脚本。用于观察回放效果和数据质量。
🚀 1. 数据收集
生成专家数据
使用 generate_expert_data.py 脚本从 Waymo 数据集中批量提取 (State, Action) 对。
# 设置 Python 路径
export PYTHONPATH=$PYTHONPATH:.:./metadrive
# 运行生成脚本
# --data_dir: Waymo 数据路径 (建议使用 exp_filtered)
# --output_dir: 结果保存路径
# --num_scenarios: 要处理的场景数量
python scripts/generate_expert_data.py \
--data_dir data/exp_filtered \
--output_dir data/training_data \
--num_scenarios 100 \
--start_index 0
生成的 .pkl 文件结构:
包含一个列表,每个元素是一条车辆轨迹(Trajectory Dictionary):
obs:(T, 45)- 观测矩阵。包含 Ego 状态 (5维) + 10辆邻居车相对信息 (40维)。acts:(T, 2)- 动作矩阵。[Steering, Accel],归一化到[-1, 1]。agent_id: 车辆 ID。scenario_id: 所属场景 ID。
内置过滤器: 脚本会自动过滤掉以下无效车辆:
- 非道路车辆:始终在停车场或路外行驶的车辆。
- 静态车辆:全称移动距离小于 5米 且速度从未超过 1m/s 的车辆(作为背景流存在,不收集数据)。
🔍 2. 数据可视化与验证
回放可视化
使用 visualize_replay.py 直观地观察回放效果,确认车辆行为是否自然,以及过滤逻辑是否生效。
# 运行可视化
# --horizon: 回放的最大步数 (Waymo 场景通常为 90 或 198 步)
python scripts/visualize_replay.py \
--data_dir data/exp_filtered \
--start_index 0 \
--num_scenarios 1 \
--horizon 200
观察要点:
- 受控车辆 (Controlled Agents):控制台会显示数量(如
Controlled agents: 2)。这些是真正产生数据的车辆。 - 背景车辆:如果在渲染图中看到其他车(通常是路边停放的),但受控数量很少,说明静态过滤生效了。
数据分析
使用 analyze_expert_data.py 查看生成数据的统计分布。
python scripts/analyze_expert_data.py --data_path data/training_data/expert_data_0_100.pkl
🧠 3. 模型训练 (Next Steps)
有了 data/training_data/ 下的专家数据后,您可以开始训练 MAGAIL 模型。
训练流程
- 加载数据:使用
dataset/expert_dataset.py中的ExpertDataset类加载.pkl数据。 - 初始化 MAGAIL:
- Generator (Policy): 接收观测
(B, 45),输出动作(B, 2)。 - Discriminator: 接收状态-动作对
(s, a),判断是专家还是生成器。
- Generator (Policy): 接收观测
- 交互采样:
- 在
MultiAgentScenarioEnv(非回放模式)中运行 Policy。 - 收集 Policy 生成的轨迹。
- 在
- 对抗更新:
- 利用专家数据和 Policy 数据训练 Discriminator。
- 利用 Discriminator 的输出作为 Reward (GAIL Reward) 训练 Policy (PPO/TRPO)。
推荐配置
- Observation: 45维 (Ego + 10 Neighbors)
- Action: 2维 Continuous (Steering, Accel)
- Horizon: 200 steps
- Batch Size: 1024+ (多智能体环境下数据量很大)
Description
Languages
Python
100%