Compare commits
3 Commits
4dbea5f0a6
...
03dee0205a
| Author | SHA1 | Date | |
|---|---|---|---|
| 03dee0205a | |||
| 21c046aef0 | |||
| 265b0eade1 |
54
.gitignore
vendored
54
.gitignore
vendored
@@ -1,3 +1,57 @@
|
||||
# 日志文件
|
||||
Env/logs/
|
||||
*.log
|
||||
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
*.so
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
|
||||
# 虚拟环境
|
||||
venv/
|
||||
env/
|
||||
ENV/
|
||||
.venv
|
||||
|
||||
# IDE
|
||||
.vscode/
|
||||
.idea/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# 数据和模型文件
|
||||
data/
|
||||
runs/
|
||||
*.pkl
|
||||
*.h5
|
||||
*.ckpt
|
||||
*.pth
|
||||
*.pt
|
||||
checkpoints/
|
||||
models/
|
||||
|
||||
# 第三方库(如果已安装)
|
||||
metadrive/
|
||||
scenarionet/
|
||||
|
||||
# 系统文件
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
|
||||
52
Algorithm/bc.py
Normal file
52
Algorithm/bc.py
Normal file
@@ -0,0 +1,52 @@
|
||||
"""
|
||||
Behavior Cloning (BC) 算法:仅包含损失与单 epoch 训练/评估逻辑。
|
||||
数据加载、环境评估、日志与保存由训练脚本 (train_bc.py) 负责。
|
||||
"""
|
||||
import torch
|
||||
|
||||
|
||||
def bc_loss(policy, states, actions):
|
||||
"""
|
||||
BC 损失:负对数似然 -E[log pi(a|s)]。
|
||||
states: (B, state_dim), actions: (B, action_dim), 均在 policy 所在 device 上。
|
||||
"""
|
||||
log_pi = policy.evaluate_log_pi(states, actions)
|
||||
return -log_pi.mean()
|
||||
|
||||
|
||||
def train_bc_epoch(policy, train_loader, optimizer, device):
|
||||
"""
|
||||
训练一个 epoch,返回平均 train loss。
|
||||
policy 与 optimizer 由调用方管理,本函数只做前向、损失、反向与 step。
|
||||
"""
|
||||
policy.train()
|
||||
total_loss = 0.0
|
||||
n_batches = 0
|
||||
for states, actions in train_loader:
|
||||
states = states.to(device)
|
||||
actions = actions.to(device)
|
||||
loss = bc_loss(policy, states, actions)
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
total_loss += loss.item()
|
||||
n_batches += 1
|
||||
return total_loss / n_batches if n_batches else 0.0
|
||||
|
||||
|
||||
def eval_bc_epoch(policy, val_loader, device):
|
||||
"""
|
||||
在验证集上评估一个 epoch,返回平均 val loss(无梯度)。
|
||||
"""
|
||||
policy.eval()
|
||||
total_loss = 0.0
|
||||
n_batches = 0
|
||||
with torch.no_grad():
|
||||
for states, actions in val_loader:
|
||||
states = states.to(device)
|
||||
actions = actions.to(device)
|
||||
log_pi = policy.evaluate_log_pi(states, actions)
|
||||
loss = -log_pi.mean().item()
|
||||
total_loss += loss
|
||||
n_batches += 1
|
||||
return total_loss / n_batches if n_batches else 0.0
|
||||
Binary file not shown.
Binary file not shown.
64
Env/bc_env.py
Normal file
64
Env/bc_env.py
Normal file
@@ -0,0 +1,64 @@
|
||||
from Env.scenario_env import MultiAgentScenarioEnv
|
||||
import numpy as np
|
||||
|
||||
class BCScenarioEnv(MultiAgentScenarioEnv):
|
||||
"""
|
||||
Environment for Behavior Cloning Evaluation.
|
||||
Uses the same 45-dim observation as ExpertReplayEnv:
|
||||
- Ego State (5): x, y, vx, vy, heading
|
||||
- Neighbors (40): 10 nearest * (rel_x, rel_y, vx, vy)
|
||||
"""
|
||||
def _get_all_obs(self):
|
||||
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||
obs_dict = {}
|
||||
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
# 1. Ego State
|
||||
ego_state = [
|
||||
vehicle.position[0], vehicle.position[1],
|
||||
vehicle.velocity[0], vehicle.velocity[1],
|
||||
vehicle.heading_theta
|
||||
]
|
||||
|
||||
# 2. Neighbors
|
||||
neighbors = []
|
||||
# Iterate through all vehicles in the engine
|
||||
candidates = []
|
||||
# Use engine.agent_manager.active_agents to find neighbors
|
||||
# Note: This includes background vehicles if they are in active_agents
|
||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
if other_id == agent_id:
|
||||
continue
|
||||
|
||||
# Check if vehicle is valid/active
|
||||
# (MetaDrive manages active_agents, so they should be active)
|
||||
|
||||
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||
if dist < 30.0:
|
||||
candidates.append((dist, other_vehicle))
|
||||
|
||||
# Sort by distance
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
|
||||
# Take top 10
|
||||
top_10 = candidates[:10]
|
||||
|
||||
neighbor_feats = []
|
||||
for _, neighbor in top_10:
|
||||
neighbor_feats.extend([
|
||||
neighbor.position[0] - vehicle.position[0], # Relative pos
|
||||
neighbor.position[1] - vehicle.position[1],
|
||||
neighbor.velocity[0], # Absolute vel
|
||||
neighbor.velocity[1]
|
||||
])
|
||||
|
||||
# Pad if < 10
|
||||
missing = 10 - len(top_10)
|
||||
if missing > 0:
|
||||
neighbor_feats.extend([0.0] * (4 * missing))
|
||||
|
||||
# Flatten
|
||||
obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||
obs_dict[agent_id] = obs
|
||||
|
||||
return obs_dict
|
||||
@@ -100,6 +100,13 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
for scenario_id in _obj_to_clean_this_frame:
|
||||
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||
|
||||
# Clear vehicles we spawned via engine.spawn_object() so _object_clean_check() passes
|
||||
ids_to_clear = [v.id for v in self.controlled_agents.values()]
|
||||
if ids_to_clear:
|
||||
self.engine.clear_objects(ids_to_clear)
|
||||
self.controlled_agents.clear()
|
||||
self.controlled_agent_ids.clear()
|
||||
|
||||
self.engine.reset()
|
||||
self.reset_sensors()
|
||||
self.engine.taskMgr.step()
|
||||
@@ -114,9 +121,6 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
self.episode_rewards = defaultdict(float)
|
||||
self.episode_lengths = defaultdict(int)
|
||||
|
||||
self.controlled_agents.clear()
|
||||
self.controlled_agent_ids.clear()
|
||||
|
||||
super().reset(seed) # 初始化场景
|
||||
self._spawn_controlled_agents()
|
||||
|
||||
@@ -190,6 +194,7 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
self.controlled_agents[agent_id].before_step(action)
|
||||
|
||||
self.engine.step()
|
||||
self.engine.after_step()
|
||||
|
||||
for agent_id in action_dict:
|
||||
if agent_id in self.controlled_agents:
|
||||
|
||||
173
README.md
173
README.md
@@ -1,98 +1,121 @@
|
||||
# MAGAIL4AutoDrive
|
||||
|
||||
> 基于多智能体生成对抗模仿学习(MAGAIL)的自动驾驶训练系统 | MetaDrive + Waymo Open Motion Dataset
|
||||
基于 **MetaDrive** 仿真器和 **Waymo Open Motion Dataset** 的自动驾驶多智能体模仿学习(MAGAIL)与行为克隆(BC)训练系统。
|
||||
|
||||
本项目利用 Waymo 真实驾驶数据,通过 MetaDrive 仿真环境构建专家回放系统,提取车辆状态与动作,用于训练多智能体模仿学习算法 (MAGAIL)。
|
||||
本项目旨在从真实的 Waymo 驾驶数据中提取专家轨迹,并通过模仿学习(Imitation Learning)训练能够适应复杂交互场景的自动驾驶策略。
|
||||
|
||||
## 📁 核心模块
|
||||
## 目录结构
|
||||
|
||||
* **`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) 对。
|
||||
|
||||
```bash
|
||||
# 设置 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
|
||||
```text
|
||||
MAGAIL4AutoDrive/
|
||||
├── Algorithm/ # 强化学习与模仿学习算法实现
|
||||
│ ├── policy.py # 基础策略网络 (MLP 等)
|
||||
│ ├── ppo.py # PPO 算法实现
|
||||
│ ├── magail.py # MAGAIL 算法核心逻辑
|
||||
│ ├── disc.py # 判别器 (Discriminator) 网络
|
||||
│ └── ...
|
||||
├── Env/ # 仿真环境封装 (MetaDrive Wrapper)
|
||||
│ ├── bc_env.py # BCScenarioEnv,45 维观测(BC/MAGAIL 共用)
|
||||
│ ├── scenario_env.py # 多智能体基础场景环境
|
||||
│ ├── expert_replay_env.py # 专家轨迹回放环境(数据生成与回放)
|
||||
│ ├── inverse_dynamics.py # 逆动力学模块 (轨迹 -> 动作)
|
||||
│ ├── simple_idm_policy.py # ConstantVelocityPolicy 占位策略
|
||||
│ └── ...
|
||||
├── dataset/ # 数据集加载器
|
||||
│ ├── loader.py # 主流水线:load_expert_pkl、MAGAILExpertDataset
|
||||
│ └── expert_dataset.py # 可选 107 维/5 维管线
|
||||
├── scripts/ # 工具脚本(数据、回放、可视化、分析)
|
||||
│ ├── generate_expert_data.py # 从 Waymo 生成专家 (obs, act) pkl
|
||||
│ ├── visualize.py # 可视化统一入口(replay / policy / trajectory)
|
||||
│ ├── analyze_expert_data.py # 数据分布分析
|
||||
│ ├── launch_tensorboard.py # 启动 TensorBoard
|
||||
│ ├── README.md # 脚本用法说明
|
||||
│ └── ...
|
||||
├── data/ # 数据目录(相对路径)
|
||||
│ ├── exp_filtered/ # Waymo 场景数据
|
||||
│ ├── training_data/ # 专家 pkl 输出(generate_expert_data)
|
||||
│ └── trajectories/ # 其他轨迹 pkl(如 expert_dataset 输出)
|
||||
├── models/ # 模型保存目录(相对路径)
|
||||
│ ├── bc/ # BC 模型 (.pt)
|
||||
│ └── magail/ # MAGAIL 模型 (*_actor.pth, *_critic.pth)
|
||||
├── logs/ # 训练日志 (TensorBoard)
|
||||
│ ├── bc/
|
||||
│ └── magail/
|
||||
├── train_bc.py # [根目录] BC 训练
|
||||
├── train_magail.py # [根目录] MAGAIL 训练
|
||||
└── README.md
|
||||
```
|
||||
|
||||
**生成的 `.pkl` 文件结构**:
|
||||
包含一个列表,每个元素是一条车辆轨迹(Trajectory Dictionary):
|
||||
* `obs`: `(T, 45)` - 观测矩阵。包含 Ego 状态 (5维) + 10辆邻居车相对信息 (40维)。
|
||||
* `acts`: `(T, 2)` - 动作矩阵。`[Steering, Accel]`,归一化到 `[-1, 1]`。
|
||||
* `agent_id`: 车辆 ID。
|
||||
* `scenario_id`: 所属场景 ID。
|
||||
## 路径约定(相对项目根)
|
||||
|
||||
**内置过滤器**:
|
||||
脚本会自动过滤掉以下无效车辆:
|
||||
1. **非道路车辆**:始终在停车场或路外行驶的车辆。
|
||||
2. **静态车辆**:全称移动距离小于 5米 且速度从未超过 1m/s 的车辆(作为背景流存在,不收集数据)。
|
||||
- **数据**:Waymo 场景 `data/exp_filtered`;专家 pkl `data/training_data`;其他轨迹 `data/trajectories`
|
||||
- **模型**:BC `models/bc/`,MAGAIL `models/magail/`
|
||||
- **日志**:TensorBoard 写入 `logs/bc/`、`logs/magail/`
|
||||
|
||||
---
|
||||
所有默认路径均为相对项目根,便于在不同设备上复用。
|
||||
|
||||
## 🔍 2. 数据可视化与验证
|
||||
## 数据处理流程
|
||||
|
||||
### 回放可视化
|
||||
使用 `visualize_replay.py` 直观地观察回放效果,确认车辆行为是否自然,以及过滤逻辑是否生效。
|
||||
从 Waymo Motion 原始数据到本项目训练用专家 pkl,依次为:
|
||||
|
||||
**1) 下载 Waymo Motion(TFRecord)**
|
||||
安装 `gsutil` 并登录 Google 账号后,例如只下载 training_20s:
|
||||
|
||||
```bash
|
||||
# 运行可视化
|
||||
# --horizon: 回放的最大步数 (Waymo 场景通常为 90 或 198 步)
|
||||
python scripts/visualize_replay.py \
|
||||
--data_dir data/exp_filtered \
|
||||
--start_index 0 \
|
||||
--num_scenarios 1 \
|
||||
--horizon 200
|
||||
gsutil -m cp -r "gs://waymo_open_dataset_motion_v_1_2_0/uncompressed/scenario/training_20s" ./waymo/
|
||||
```
|
||||
|
||||
**观察要点**:
|
||||
* **受控车辆 (Controlled Agents)**:控制台会显示数量(如 `Controlled agents: 2`)。这些是真正产生数据的车辆。
|
||||
* **背景车辆**:如果在渲染图中看到其他车(通常是路边停放的),但受控数量很少,说明静态过滤生效了。
|
||||
|
||||
### 数据分析
|
||||
使用 `analyze_expert_data.py` 查看生成数据的统计分布。
|
||||
**2) ScenarioNet Convert(TFRecord → ScenarioNet 场景库)**
|
||||
需安装 ScenarioNet、MetaDrive 及 TensorFlow 2.11、protobuf 3.20;转换时不用 GPU。
|
||||
|
||||
```bash
|
||||
python scripts/analyze_expert_data.py --data_path data/training_data/expert_data_0_100.pkl
|
||||
python -m scenarionet.convert_waymo -d data/exp_converted --raw_data_path ./waymo/training_20s --num_workers 64
|
||||
```
|
||||
|
||||
---
|
||||
**3) ScenarioNet Filter(按需筛选场景)**
|
||||
从 convert 得到的场景库中筛掉含红绿灯、天桥等场景,输出到如 `data/exp_filtered`。具体命令以 ScenarioNet 文档为准(Operations → Filter)。
|
||||
|
||||
## 🧠 3. 模型训练 (Next Steps)
|
||||
**4) 本项目:生成专家 pkl**
|
||||
使用筛选后的场景目录,生成训练用 pkl 到 `data/training_data`:
|
||||
|
||||
有了 `data/training_data/` 下的专家数据后,您可以开始训练 MAGAIL 模型。
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||
```
|
||||
|
||||
### 训练流程
|
||||
1. **加载数据**:使用 `dataset/expert_dataset.py` 中的 `ExpertDataset` 类加载 `.pkl` 数据。
|
||||
2. **初始化 MAGAIL**:
|
||||
* **Generator (Policy)**: 接收观测 `(B, 45)`,输出动作 `(B, 2)`。
|
||||
* **Discriminator**: 接收状态-动作对 `(s, a)`,判断是专家还是生成器。
|
||||
3. **交互采样**:
|
||||
* 在 `MultiAgentScenarioEnv`(非回放模式)中运行 Policy。
|
||||
* 收集 Policy 生成的轨迹。
|
||||
4. **对抗更新**:
|
||||
* 利用专家数据和 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+ (多智能体环境下数据量很大)
|
||||
### 1. 数据准备
|
||||
使用 `scripts/generate_expert_data.py` 将 Waymo 数据转换为训练用 `.pkl`,输出到 `data/training_data/`。
|
||||
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100
|
||||
```
|
||||
|
||||
### 2. 行为克隆 (BC)
|
||||
- **训练**:`python train_bc.py`(模型保存到 `models/bc/`,日志到 `logs/bc/`)
|
||||
- **可视化**:`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt`
|
||||
|
||||
### 3. 多智能体对抗模仿学习 (MAGAIL)
|
||||
- **训练**:`python train_magail.py`(模型保存到 `models/magail/`,日志到 `logs/magail/`)
|
||||
- **可视化**:`python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth`
|
||||
|
||||
### 4. 可视化统一入口
|
||||
可视化统一使用 `scripts/visualize.py`,子命令:`replay`(场景回放)、`policy`(BC/MAGAIL 策略)、`trajectory`(专家轨迹 2D 动画)。详见 [scripts/README.md](scripts/README.md)。
|
||||
|
||||
## 文件与模块职责
|
||||
|
||||
### 根目录脚本
|
||||
- **train_bc.py**:BC 训练,从 `dataset.loader` 加载专家 pkl,模型与日志写入 `models/bc/`、`logs/bc/`
|
||||
- **train_magail.py**:MAGAIL 训练,环境使用 `BCScenarioEnv`(45 维),从 `dataset.loader` 加载专家数据,模型与日志写入 `models/magail/`、`logs/magail/`
|
||||
|
||||
### Env 模块
|
||||
- **Env/bc_env.py**:`BCScenarioEnv`,45 维观测(Ego 5 维 + 10 邻居×4 维),BC 与 MAGAIL 训练/评估共用
|
||||
- **Env/scenario_env.py**:`MultiAgentScenarioEnv` 基类,Waymo 场景加载与步进
|
||||
- **Env/expert_replay_env.py**:专家轨迹回放与逆动力学动作,供 `generate_expert_data.py` 与回放可视化
|
||||
- **Env/inverse_dynamics.py**:轨迹 → 油门/转向动作
|
||||
|
||||
### Algorithm 模块
|
||||
- **Algorithm/policy.py**:`StateIndependentPolicy`,BC 使用的 MLP 策略
|
||||
|
||||
### scripts 目录
|
||||
工具脚本用途与用法见 [scripts/README.md](scripts/README.md)。
|
||||
|
||||
@@ -244,6 +244,7 @@ class ExpertTrajectoryDataset(Dataset):
|
||||
print(f" 观测维度: {obs_dim} (应为107)")
|
||||
|
||||
if save_path:
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
with open(save_path, "wb") as f:
|
||||
pickle.dump({
|
||||
"trajectories": all_trajectories,
|
||||
@@ -282,7 +283,7 @@ if __name__ == "__main__":
|
||||
trajectories, observations = ExpertTrajectoryDataset.collect_with_full_obs(
|
||||
env_config,
|
||||
num_scenarios=10,
|
||||
save_path="./expert_trajectories_full.pkl"
|
||||
save_path="data/trajectories/expert_trajectories_full.pkl"
|
||||
)
|
||||
|
||||
if len(trajectories) > 0:
|
||||
|
||||
103
dataset/loader.py
Normal file
103
dataset/loader.py
Normal file
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
统一数据加载:BC/MAGAIL 训练用专家 pkl 的加载函数与 Dataset。
|
||||
主训练流水线使用本模块;dataset/expert_dataset.py 为可选 107 维/5 维管线。
|
||||
"""
|
||||
import os
|
||||
import glob
|
||||
import pickle
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
def load_expert_pkl(expert_data_path):
|
||||
"""从目录或单个 pkl 加载专家 (obs, acts),返回 concat 后的 obs_data, act_data。"""
|
||||
if os.path.isdir(expert_data_path):
|
||||
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
||||
if not pkl_files:
|
||||
raise FileNotFoundError(f"No .pkl files in {expert_data_path}")
|
||||
print(f"Found {len(pkl_files)} pickle files in {expert_data_path}")
|
||||
elif os.path.exists(expert_data_path):
|
||||
pkl_files = [expert_data_path]
|
||||
else:
|
||||
raise FileNotFoundError(f"Expert data path not found: {expert_data_path}")
|
||||
|
||||
obs_data, act_data = [], []
|
||||
for pkl_file in pkl_files:
|
||||
try:
|
||||
with open(pkl_file, "rb") as f:
|
||||
data = pickle.load(f)
|
||||
if isinstance(data, list):
|
||||
for traj in data:
|
||||
if "obs" in traj and "acts" in traj:
|
||||
obs_data.append(traj["obs"])
|
||||
act_data.append(traj["acts"])
|
||||
elif isinstance(data, dict):
|
||||
if "observations" in data and "actions" in data:
|
||||
obs_data.append(data["observations"])
|
||||
act_data.append(data["actions"])
|
||||
else:
|
||||
print(f"Skipping {pkl_file}: Unknown data format {type(data)}")
|
||||
except Exception as e:
|
||||
print(f"Error loading {pkl_file}: {e}")
|
||||
|
||||
if len(obs_data) == 0:
|
||||
raise ValueError("No valid data loaded from provided path.")
|
||||
obs_data = np.concatenate(obs_data, axis=0)
|
||||
act_data = np.concatenate(act_data, axis=0)
|
||||
print(f"Total loaded samples: {len(obs_data)}")
|
||||
return obs_data, act_data
|
||||
|
||||
|
||||
class MAGAILExpertDataset(Dataset):
|
||||
def __init__(self, data_dir, transform=None):
|
||||
"""
|
||||
Args:
|
||||
data_dir (str): Directory containing .pkl files from generate_expert_data.py
|
||||
transform (callable, optional): Optional transform to be applied on a sample.
|
||||
"""
|
||||
self.data_dir = data_dir
|
||||
self.transform = transform
|
||||
self.trajectories = []
|
||||
self.flat_data = [] # (obs, act) pairs
|
||||
|
||||
# Load all .pkl files
|
||||
pkl_files = glob.glob(os.path.join(data_dir, "*.pkl"))
|
||||
print(f"Loading data from {len(pkl_files)} files in {data_dir}...")
|
||||
|
||||
for pkl_file in pkl_files:
|
||||
try:
|
||||
with open(pkl_file, "rb") as f:
|
||||
data = pickle.load(f)
|
||||
# data is a list of dicts: {'obs': (T, 45), 'acts': (T, 2), ...}
|
||||
self.trajectories.extend(data)
|
||||
except Exception as e:
|
||||
print(f"Error loading {pkl_file}: {e}")
|
||||
|
||||
# Flatten for training Discriminator/BC
|
||||
print(f"Processing {len(self.trajectories)} trajectories...")
|
||||
for traj in self.trajectories:
|
||||
obs = traj["obs"]
|
||||
acts = traj["acts"]
|
||||
|
||||
# obs: (T, 45), acts: (T, 2)
|
||||
for i in range(len(obs)):
|
||||
self.flat_data.append((obs[i], acts[i]))
|
||||
|
||||
print(f"Total samples: {len(self.flat_data)}")
|
||||
|
||||
def __len__(self):
|
||||
return len(self.flat_data)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
obs, act = self.flat_data[idx]
|
||||
|
||||
obs = torch.from_numpy(obs).float()
|
||||
act = torch.from_numpy(act).float()
|
||||
|
||||
sample = {"state": obs, "action": act}
|
||||
|
||||
if self.transform:
|
||||
sample = self.transform(sample)
|
||||
|
||||
return sample
|
||||
@@ -1,61 +0,0 @@
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
import pickle
|
||||
import numpy as np
|
||||
import os
|
||||
import glob
|
||||
|
||||
class MAGAILExpertDataset(Dataset):
|
||||
def __init__(self, data_dir, transform=None):
|
||||
"""
|
||||
Args:
|
||||
data_dir (str): Directory containing .pkl files from generate_expert_data.py
|
||||
transform (callable, optional): Optional transform to be applied on a sample.
|
||||
"""
|
||||
self.data_dir = data_dir
|
||||
self.transform = transform
|
||||
self.trajectories = []
|
||||
self.flat_data = [] # (obs, act) pairs
|
||||
|
||||
# Load all .pkl files
|
||||
pkl_files = glob.glob(os.path.join(data_dir, "*.pkl"))
|
||||
print(f"Loading data from {len(pkl_files)} files in {data_dir}...")
|
||||
|
||||
for pkl_file in pkl_files:
|
||||
try:
|
||||
with open(pkl_file, 'rb') as f:
|
||||
data = pickle.load(f)
|
||||
# data is a list of dicts: {'obs': (T, 45), 'acts': (T, 2), ...}
|
||||
self.trajectories.extend(data)
|
||||
except Exception as e:
|
||||
print(f"Error loading {pkl_file}: {e}")
|
||||
|
||||
# Flatten for training Discriminator/BC
|
||||
print(f"Processing {len(self.trajectories)} trajectories...")
|
||||
for traj in self.trajectories:
|
||||
obs = traj['obs']
|
||||
acts = traj['acts']
|
||||
|
||||
# obs: (T, 45), acts: (T, 2)
|
||||
# We pair them up
|
||||
for i in range(len(obs)):
|
||||
self.flat_data.append((obs[i], acts[i]))
|
||||
|
||||
print(f"Total samples: {len(self.flat_data)}")
|
||||
|
||||
def __len__(self):
|
||||
return len(self.flat_data)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
obs, act = self.flat_data[idx]
|
||||
|
||||
# Convert to tensor
|
||||
obs = torch.from_numpy(obs).float()
|
||||
act = torch.from_numpy(act).float()
|
||||
|
||||
sample = {'state': obs, 'action': act}
|
||||
|
||||
if self.transform:
|
||||
sample = self.transform(sample)
|
||||
|
||||
return sample
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
76
scripts/README.md
Normal file
76
scripts/README.md
Normal file
@@ -0,0 +1,76 @@
|
||||
# scripts 工具脚本说明
|
||||
|
||||
本目录包含数据生成、回放、可视化与分析等工具脚本。训练脚本(`train_bc.py`、`train_magail.py`)位于项目根目录。
|
||||
|
||||
## 路径约定(相对项目根)
|
||||
|
||||
- **数据**:`data/exp_filtered`(Waymo 场景)、`data/training_data`(专家 pkl 输出)
|
||||
- **模型**:`models/bc/`(BC)、`models/magail/`(MAGAIL)
|
||||
- **日志**:`logs/bc/`、`logs/magail/`(TensorBoard)
|
||||
|
||||
---
|
||||
|
||||
## 脚本列表与用法
|
||||
|
||||
### 数据生成
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [generate_expert_data.py](generate_expert_data.py) | 从 Waymo 数据生成专家 (obs, act) 的 pkl | `python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100` |
|
||||
|
||||
**常用参数**:`--data_dir`(默认 `data/exp_filtered`)、`--output_dir`(默认 `data/training_data`)、`--start_index`、`--num_scenarios`。
|
||||
|
||||
---
|
||||
|
||||
### 可视化(统一入口)
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [visualize.py](visualize.py) | **replay**:场景回放(ExpertReplayEnv);**policy**:BC/MAGAIL 策略;**trajectory**:专家轨迹 2D 动画 | 见下方 |
|
||||
|
||||
**子命令**:
|
||||
|
||||
- **replay**(原始专家轨迹回放):
|
||||
```bash
|
||||
python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500
|
||||
```
|
||||
|
||||
- **policy**(BC 或 MAGAIL 训练策略):
|
||||
```bash
|
||||
python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1
|
||||
python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth --num_scenarios 1 --deterministic
|
||||
```
|
||||
|
||||
- **trajectory**(专家轨迹 matplotlib 俯视图动画):
|
||||
```bash
|
||||
python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_idx 0
|
||||
```
|
||||
|
||||
**公共参数**:`--data_dir`(默认 `data/exp_filtered`)、`--start_index`、`--num_scenarios`、`--horizon`。policy 模式另有 `--policy_type`(auto/bc/magail)、`--model_path`、`--deterministic`(仅 MAGAIL)。
|
||||
|
||||
---
|
||||
|
||||
### 数据分析与检查
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [analyze_expert_data.py](analyze_expert_data.py) | 分析专家数据分布与统计 | 见脚本内 `__main__`(依赖 env 与数据目录配置) |
|
||||
| [check_track_fields.py](check_track_fields.py) | 检查 Waymo 轨迹字段 | 见脚本内 `__main__` |
|
||||
| [check_database_info.py](check_database_info.py) | 检查数据库/场景信息 | 见脚本内 `__main__`(含硬编码路径,可按需改为 `data/exp_filtered`) |
|
||||
|
||||
---
|
||||
|
||||
### 其他
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [launch_tensorboard.py](launch_tensorboard.py) | 启动 TensorBoard | `python scripts/launch_tensorboard.py --logdir logs`(或 `logs/bc` / `logs/magail`) |
|
||||
|
||||
---
|
||||
|
||||
## 与训练流程的对应关系
|
||||
|
||||
1. **数据准备**:`generate_expert_data.py` → 输出到 `data/training_data/*.pkl`
|
||||
2. **BC 训练**:根目录 `train_bc.py` → 模型保存到 `models/bc/`,日志到 `logs/bc/`
|
||||
3. **MAGAIL 训练**:根目录 `train_magail.py` → 模型保存到 `models/magail/`,日志到 `logs/magail/`
|
||||
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
||||
@@ -153,8 +153,8 @@ def generate_data(args):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/exp_filtered", help="Path to Waymo pickles (or filtered index)")
|
||||
parser.add_argument("--output_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/training", help="Output directory")
|
||||
parser.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Path to Waymo pickles (or filtered index)")
|
||||
parser.add_argument("--output_dir", type=str, default="data/training_data", help="Output directory")
|
||||
parser.add_argument("--start_index", type=int, default=0)
|
||||
parser.add_argument("--num_scenarios", type=int, default=10)
|
||||
|
||||
|
||||
18
scripts/launch_tensorboard.py
Normal file
18
scripts/launch_tensorboard.py
Normal file
@@ -0,0 +1,18 @@
|
||||
import sys
|
||||
import types
|
||||
import os
|
||||
|
||||
# Mock imghdr module for Python 3.13 compatibility
|
||||
# TensorBoard depends on imghdr which was removed in Python 3.13
|
||||
if sys.version_info >= (3, 13):
|
||||
if 'imghdr' not in sys.modules:
|
||||
imghdr_mock = types.ModuleType('imghdr')
|
||||
imghdr_mock.what = lambda filename, h=None: None
|
||||
# Mock tests list which tensorboard appends to
|
||||
imghdr_mock.tests = []
|
||||
sys.modules['imghdr'] = imghdr_mock
|
||||
|
||||
from tensorboard import main as tb_main
|
||||
|
||||
if __name__ == '__main__':
|
||||
sys.exit(tb_main.run_main())
|
||||
395
scripts/visualize.py
Normal file
395
scripts/visualize.py
Normal file
@@ -0,0 +1,395 @@
|
||||
"""
|
||||
Unified visualization: replay (scenario replay), policy (BC/MAGAIL), trajectory (2D expert trajectory animation).
|
||||
Usage: python scripts/visualize.py <replay|policy|trajectory> [args...]
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
# --- Replay ---
|
||||
def _run_replay(args):
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
data_path = os.path.abspath(args.data_dir)
|
||||
if not os.path.exists(data_path):
|
||||
raise ValueError(f"Data directory {data_path} not found")
|
||||
|
||||
from metadrive.scenario.utils import read_dataset_summary
|
||||
_, summary_lookup, _ = read_dataset_summary(data_path)
|
||||
if args.start_index >= len(summary_lookup):
|
||||
raise ValueError(
|
||||
f"start_index={args.start_index} out of range. Dataset has {len(summary_lookup)} scenarios."
|
||||
)
|
||||
max_available = len(summary_lookup) - args.start_index
|
||||
num_to_run = min(args.num_scenarios, max_available)
|
||||
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": -1,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + num_to_run):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
||||
|
||||
for step in range(args.horizon):
|
||||
obs, rewards, dones, infos = env.step(None)
|
||||
env.render(
|
||||
mode="top_down",
|
||||
text={"Step": step, "Agents": len(env.controlled_agents), "Scenario": i},
|
||||
)
|
||||
time.sleep(0.05)
|
||||
if dones["__all__"]:
|
||||
print(f"Scenario {i} finished at step {step}")
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted by user")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
env.close()
|
||||
print("Environment closed.")
|
||||
|
||||
|
||||
# --- Policy (BC / MAGAIL) ---
|
||||
def _resolve_data_dir(data_dir_arg):
|
||||
if data_dir_arg:
|
||||
data_dir = data_dir_arg
|
||||
else:
|
||||
data_dir = os.path.join(project_root, "data", "exp_filtered")
|
||||
if not os.path.exists(data_dir):
|
||||
data_dir = os.path.join(project_root, "data", "exp_converted")
|
||||
if not os.path.exists(data_dir):
|
||||
raise FileNotFoundError(f"Data directory not found at {data_dir}. Please specify --data_dir.")
|
||||
return data_dir
|
||||
|
||||
|
||||
def _resolve_model_path(model_path, policy_type):
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
if policy_type == "bc":
|
||||
candidate = os.path.join(project_root, "models", "bc", os.path.basename(model_path))
|
||||
else:
|
||||
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path))
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
if policy_type == "magail" and not model_path.endswith("_actor.pth"):
|
||||
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path) + "_actor.pth")
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
raise FileNotFoundError(f"Model path {model_path} not found.")
|
||||
|
||||
|
||||
def _run_policy(args):
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
from metadrive.engine.engine_utils import close_engine
|
||||
|
||||
policy_type = (args.policy_type or "auto").lower()
|
||||
if policy_type == "auto":
|
||||
policy_type = "bc" if args.model_path.endswith(".pt") else "magail"
|
||||
|
||||
data_dir = _resolve_data_dir(args.data_dir)
|
||||
data_path = os.path.abspath(data_dir)
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True,
|
||||
"sequential_seed": True,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": args.num_scenarios,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
||||
try:
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
except Exception as e:
|
||||
print(f"Error init env: {e}. Trying to close lingering engine...")
|
||||
try:
|
||||
close_engine()
|
||||
except Exception:
|
||||
pass
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
|
||||
state_dim = 45
|
||||
action_dim = 2
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
model_path = _resolve_model_path(args.model_path, policy_type)
|
||||
print(f"Loading model from {model_path}...")
|
||||
|
||||
if policy_type == "bc":
|
||||
from Algorithm.policy import StateIndependentPolicy
|
||||
policy = StateIndependentPolicy(
|
||||
state_shape=(state_dim,),
|
||||
action_shape=(action_dim,),
|
||||
hidden_units=(256, 256),
|
||||
hidden_activation=torch.nn.Tanh(),
|
||||
).to(device)
|
||||
policy.load_state_dict(torch.load(model_path, map_location=device))
|
||||
policy.eval()
|
||||
else:
|
||||
from train_magail import Actor
|
||||
actor = Actor(state_dim, action_dim).to(device)
|
||||
actor.load_state_dict(torch.load(model_path, map_location=device))
|
||||
actor.eval()
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + args.num_scenarios):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs_dict = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting {i}: {e}. Skipping.")
|
||||
try:
|
||||
close_engine()
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(obs_dict)}")
|
||||
step_count = 0
|
||||
episode_reward = 0.0
|
||||
|
||||
while True:
|
||||
agent_ids = list(obs_dict.keys())
|
||||
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
if policy_type == "bc":
|
||||
actions_np = policy(obs_tensor).cpu().numpy()
|
||||
else:
|
||||
dist = actor(obs_tensor)
|
||||
if args.deterministic:
|
||||
actions_np = torch.tanh(dist.mean).cpu().numpy()
|
||||
else:
|
||||
actions_np = torch.tanh(dist.sample()).cpu().numpy()
|
||||
|
||||
actions = {aid: actions_np[idx].flatten() for idx, aid in enumerate(agent_ids)}
|
||||
obs_dict, rewards, dones, infos = env.step(actions)
|
||||
episode_reward += sum(rewards.values())
|
||||
|
||||
env.render(
|
||||
mode="top_down",
|
||||
text={
|
||||
"Scenario": i,
|
||||
"Step": step_count,
|
||||
"Agents": len(obs_dict),
|
||||
"Total Reward": f"{episode_reward:.2f}",
|
||||
},
|
||||
)
|
||||
step_count += 1
|
||||
|
||||
if dones["__all__"] or step_count >= args.horizon:
|
||||
print(f"Scenario finished at step {step_count}, reward {episode_reward:.2f}")
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted.")
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
# --- Trajectory (matplotlib 2D animation) ---
|
||||
def _build_expert_trajectories_from_env(env):
|
||||
"""Build expert_trajectories dict from env (ExpertReplayEnv has traffic_manager.current_traffic_data)."""
|
||||
if hasattr(env, "expert_trajectories") and env.expert_trajectories:
|
||||
return env.expert_trajectories
|
||||
if not hasattr(env, "engine") or not hasattr(env.engine, "traffic_manager"):
|
||||
return {}
|
||||
from metadrive.type import MetaDriveType
|
||||
data = getattr(env.engine.traffic_manager, "current_traffic_data", None)
|
||||
if not data:
|
||||
return {}
|
||||
expert_trajs = {}
|
||||
for scenario_id, track in data.items():
|
||||
if track.get("type") != MetaDriveType.VEHICLE or "state" not in track:
|
||||
continue
|
||||
state = track["state"]
|
||||
positions = state.get("position")
|
||||
if positions is None:
|
||||
continue
|
||||
valid = state.get("valid", np.ones(len(positions), dtype=bool))
|
||||
valid = np.asarray(valid).flatten()
|
||||
if valid.size != len(positions):
|
||||
valid = np.ones(len(positions), dtype=bool)
|
||||
first_show = int(np.argmax(valid)) if valid.any() else 0
|
||||
last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else len(positions) - 1
|
||||
obj_id = track.get("metadata", {}).get("object_id", str(scenario_id))
|
||||
expert_trajs[obj_id] = {
|
||||
"positions": np.asarray(positions),
|
||||
"start_timestep": first_show,
|
||||
"end_timestep": last_show,
|
||||
}
|
||||
return expert_trajs
|
||||
|
||||
|
||||
def _run_trajectory_animation(expert_trajs, scenario_idx):
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.animation import FuncAnimation
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
print("No expert trajectories to visualize.")
|
||||
return
|
||||
|
||||
fig, ax = plt.subplots(figsize=(12, 12))
|
||||
max_timestep = max(t["end_timestep"] for t in expert_trajs.values())
|
||||
min_timestep = min(t["start_timestep"] for t in expert_trajs.values())
|
||||
|
||||
colors = plt.cm.tab10(np.linspace(0, 1, len(expert_trajs)))
|
||||
for idx, (obj_id, traj) in enumerate(expert_trajs.items()):
|
||||
positions = np.asarray(traj["positions"])
|
||||
if positions.ndim >= 2:
|
||||
positions = positions[:, :2]
|
||||
else:
|
||||
continue
|
||||
ax.plot(
|
||||
positions[:, 0], positions[:, 1],
|
||||
color=colors[idx], alpha=0.3, linewidth=1,
|
||||
label=f"Vehicle {str(obj_id)[:6]}",
|
||||
)
|
||||
|
||||
scatter = ax.scatter([], [], s=200, c="red", marker="o", edgecolors="black", linewidths=2)
|
||||
time_text = ax.text(0.02, 0.95, "", transform=ax.transAxes, fontsize=14)
|
||||
ax.set_xlabel("X (m)")
|
||||
ax.set_ylabel("Y (m)")
|
||||
ax.set_title(f"Expert Trajectory Visualization - Scenario {scenario_idx}")
|
||||
ax.legend(loc="upper right", fontsize=8)
|
||||
ax.grid(True, alpha=0.3)
|
||||
ax.axis("equal")
|
||||
|
||||
def update(frame):
|
||||
current_time = min_timestep + frame
|
||||
current_positions = []
|
||||
for traj in expert_trajs.values():
|
||||
st, et = traj["start_timestep"], traj["end_timestep"]
|
||||
if st <= current_time <= et:
|
||||
pos = np.asarray(traj["positions"])
|
||||
if pos.ndim >= 2:
|
||||
pos = pos[current_time - st, :2]
|
||||
else:
|
||||
continue
|
||||
current_positions.append(pos)
|
||||
if current_positions:
|
||||
scatter.set_offsets(np.array(current_positions))
|
||||
time_text.set_text(f"Time: {frame * 0.1:.1f}s (Frame {frame})")
|
||||
return scatter, time_text
|
||||
|
||||
anim = FuncAnimation(
|
||||
fig, update, frames=max_timestep - min_timestep + 1,
|
||||
interval=100, blit=True, repeat=True,
|
||||
)
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
return anim
|
||||
|
||||
|
||||
def _run_trajectory(args):
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
data_dir = _resolve_data_dir(args.data_dir)
|
||||
data_path = os.path.abspath(data_dir)
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": 500,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.scenario_idx,
|
||||
"num_scenarios": 1,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
try:
|
||||
env.reset(seed=args.scenario_idx)
|
||||
expert_trajs = _build_expert_trajectories_from_env(env)
|
||||
_run_trajectory_animation(expert_trajs, args.scenario_idx)
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
# --- Main ---
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Unified visualization: replay, policy (BC/MAGAIL), trajectory.",
|
||||
)
|
||||
subparsers = parser.add_subparsers(dest="mode", required=True, help="replay | policy | trajectory")
|
||||
|
||||
# Common args for data_dir (used by all)
|
||||
def add_common_data_args(p):
|
||||
p.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Waymo scenario directory")
|
||||
p.add_argument("--start_index", type=int, default=0)
|
||||
p.add_argument("--num_scenarios", type=int, default=1)
|
||||
p.add_argument("--horizon", type=int, default=200)
|
||||
|
||||
# replay
|
||||
pr = subparsers.add_parser("replay", help="Replay scenario with ExpertReplayEnv (no policy)")
|
||||
add_common_data_args(pr)
|
||||
pr.set_defaults(horizon=500)
|
||||
|
||||
# policy
|
||||
pp = subparsers.add_parser("policy", help="Visualize BC or MAGAIL trained policy")
|
||||
add_common_data_args(pp)
|
||||
pp.add_argument("--policy_type", type=str, default="auto", choices=["auto", "bc", "magail"])
|
||||
pp.add_argument("--model_path", type=str, default="models/bc/policy_best.pt")
|
||||
pp.add_argument("--deterministic", action="store_true", help="MAGAIL: use mean action")
|
||||
|
||||
# trajectory
|
||||
pt = subparsers.add_parser("trajectory", help="2D matplotlib animation of expert trajectories")
|
||||
pt.add_argument("--data_dir", type=str, default="data/exp_filtered")
|
||||
pt.add_argument("--scenario_idx", type=int, default=0)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Resolve data_dir relative to project root when default
|
||||
if args.mode != "trajectory":
|
||||
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||
else:
|
||||
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||
|
||||
if args.mode == "replay":
|
||||
_run_replay(args)
|
||||
elif args.mode == "policy":
|
||||
_run_policy(args)
|
||||
elif args.mode == "trajectory":
|
||||
_run_trajectory(args)
|
||||
else:
|
||||
parser.error(f"Unknown mode: {args.mode}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,105 +0,0 @@
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 添加路径
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
project_root = os.path.dirname(current_dir)
|
||||
env_dir = os.path.join(project_root, "Env")
|
||||
sys.path.insert(0, project_root)
|
||||
sys.path.insert(0, env_dir)
|
||||
|
||||
# 现在可以导入了
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.animation import FuncAnimation
|
||||
|
||||
class DummyPolicy:
|
||||
"""
|
||||
占位策略,用于数据检查时初始化环境
|
||||
不需要实际执行动作,只是为了满足环境初始化要求
|
||||
"""
|
||||
def act(self, *args, **kwargs):
|
||||
# 返回零动作 [throttle, steering]
|
||||
return np.array([0.0, 0.0])
|
||||
|
||||
def visualize_expert_trajectory(env, scenario_idx=0):
|
||||
"""
|
||||
可视化专家轨迹的俯视图动画
|
||||
"""
|
||||
env.reset()
|
||||
expert_trajs = env.expert_trajectories
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
print("当前场景无专家轨迹")
|
||||
return
|
||||
|
||||
# 设置绘图
|
||||
fig, ax = plt.subplots(figsize=(12, 12))
|
||||
|
||||
# 获取所有轨迹的最大时间长度
|
||||
max_timestep = max(traj["end_timestep"] for traj in expert_trajs.values())
|
||||
min_timestep = min(traj["start_timestep"] for traj in expert_trajs.values())
|
||||
|
||||
# 绘制完整轨迹(淡色)
|
||||
colors = plt.cm.tab10(np.linspace(0, 1, len(expert_trajs)))
|
||||
for idx, (obj_id, traj) in enumerate(expert_trajs.items()):
|
||||
positions = traj["positions"][:, :2]
|
||||
ax.plot(positions[:, 0], positions[:, 1],
|
||||
color=colors[idx], alpha=0.3, linewidth=1,
|
||||
label=f'Vehicle {obj_id[:6]}')
|
||||
|
||||
# 初始化当前位置标记
|
||||
scatter = ax.scatter([], [], s=200, c='red', marker='o', edgecolors='black', linewidths=2)
|
||||
time_text = ax.text(0.02, 0.95, '', transform=ax.transAxes, fontsize=14)
|
||||
|
||||
ax.set_xlabel('X (m)')
|
||||
ax.set_ylabel('Y (m)')
|
||||
ax.set_title(f'Expert Trajectory Visualization - Scenario {scenario_idx}')
|
||||
ax.legend(loc='upper right', fontsize=8)
|
||||
ax.grid(True, alpha=0.3)
|
||||
ax.axis('equal')
|
||||
|
||||
def update(frame):
|
||||
current_time = min_timestep + frame
|
||||
|
||||
# 收集当前时间所有车辆的位置
|
||||
current_positions = []
|
||||
for traj in expert_trajs.values():
|
||||
if traj["start_timestep"] <= current_time <= traj["end_timestep"]:
|
||||
idx = current_time - traj["start_timestep"]
|
||||
pos = traj["positions"][idx, :2]
|
||||
current_positions.append(pos)
|
||||
|
||||
if len(current_positions) > 0:
|
||||
current_positions = np.array(current_positions)
|
||||
scatter.set_offsets(current_positions)
|
||||
|
||||
time_text.set_text(f'Time: {frame * 0.1:.1f}s (Frame {frame})')
|
||||
return scatter, time_text
|
||||
|
||||
anim = FuncAnimation(fig, update, frames=max_timestep-min_timestep+1,
|
||||
interval=100, blit=True, repeat=True)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
|
||||
return anim
|
||||
|
||||
if __name__ == "__main__":
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/mdsn"
|
||||
data_dir = AssetLoader.file_path(WAYMO_DATA_DIR, "exp_filtered", unix_style=False)
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
},
|
||||
agent2policy=DummyPolicy()
|
||||
)
|
||||
|
||||
# 可视化第一个场景
|
||||
anim = visualize_expert_trajectory(env, scenario_idx=0)
|
||||
@@ -1,93 +0,0 @@
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
# Add project root to Python path so we can import Env module
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
def visualize_replay(args):
|
||||
data_path = os.path.abspath(args.data_dir)
|
||||
if not os.path.exists(data_path):
|
||||
raise ValueError(f"Data directory {data_path} not found")
|
||||
|
||||
# Same as data generation: avoid MetaDrive assertion when requested num_scenarios > available.
|
||||
from metadrive.scenario.utils import read_dataset_summary
|
||||
_, summary_lookup, _ = read_dataset_summary(data_path)
|
||||
if args.start_index >= len(summary_lookup):
|
||||
raise ValueError(
|
||||
f"start_index={args.start_index} out of range. Dataset has {len(summary_lookup)} scenarios."
|
||||
)
|
||||
max_available = len(summary_lookup) - args.start_index
|
||||
num_to_run = min(args.num_scenarios, max_available)
|
||||
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True, # Enable rendering
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": -1,
|
||||
"log_level": 40, # ERROR
|
||||
# "pstats": True, # For performance debugging
|
||||
}
|
||||
|
||||
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + num_to_run):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
||||
|
||||
for step in range(args.horizon):
|
||||
# Step
|
||||
obs, rewards, dones, infos = env.step(None)
|
||||
|
||||
# Render
|
||||
env.render(mode="top_down",
|
||||
text={
|
||||
"Step": step,
|
||||
"Agents": len(env.controlled_agents),
|
||||
"Scenario": i
|
||||
})
|
||||
|
||||
# Sleep to control playback speed
|
||||
time.sleep(0.05)
|
||||
|
||||
if dones["__all__"]:
|
||||
print(f"Scenario {i} finished at step {step}")
|
||||
break
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted by user")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
env.close()
|
||||
print("Environment closed.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/exp_filtered", help="Path to Waymo data")
|
||||
parser.add_argument("--start_index", type=int, default=0)
|
||||
parser.add_argument("--num_scenarios", type=int, default=1)
|
||||
parser.add_argument("--horizon", type=int, default=500)
|
||||
|
||||
args = parser.parse_args()
|
||||
visualize_replay(args)
|
||||
146
train_bc.py
Normal file
146
train_bc.py
Normal file
@@ -0,0 +1,146 @@
|
||||
"""
|
||||
BC 训练脚本:负责数据加载、环境评估、日志与保存;BC 算法由 Algorithm.bc 提供。
|
||||
使用方式不变:python train_bc.py [--expert_data_path data/training_data] [--save_dir models/bc] ...
|
||||
"""
|
||||
import os
|
||||
import numpy as np
|
||||
import torch
|
||||
import argparse
|
||||
from torch.utils.data import DataLoader, TensorDataset
|
||||
from torch.optim import Adam
|
||||
from torch.optim.lr_scheduler import ExponentialLR
|
||||
from datetime import datetime
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from Algorithm.policy import StateIndependentPolicy
|
||||
from Algorithm.bc import train_bc_epoch, eval_bc_epoch
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
from dataset.loader import load_expert_pkl
|
||||
|
||||
|
||||
def evaluate_policy(policy, args, device):
|
||||
"""在 BCScenarioEnv 中评估策略,跑若干 episode,返回平均 reward。"""
|
||||
waymo_data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
|
||||
data_dir = os.path.join(waymo_data_dir, "exp_filtered")
|
||||
if not os.path.exists(data_dir):
|
||||
data_dir = os.path.join(waymo_data_dir, "exp_converted")
|
||||
if not os.path.exists(data_dir):
|
||||
print(f"[ERROR] Could not find scenario data in {waymo_data_dir}. Evaluation skipped.")
|
||||
return 0.0
|
||||
|
||||
env_config = {
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"horizon": 200,
|
||||
}
|
||||
env = BCScenarioEnv(env_config, agent2policy=None)
|
||||
total_rewards = []
|
||||
|
||||
try:
|
||||
for i in range(3):
|
||||
obs_dict = env.reset(seed=i)
|
||||
episode_reward = 0
|
||||
dones = {"__all__": False}
|
||||
step_count = 0
|
||||
horizon = 200
|
||||
while not dones["__all__"]:
|
||||
step_count += 1
|
||||
if step_count >= horizon:
|
||||
break
|
||||
if not obs_dict:
|
||||
obs_dict, _, dones, _ = env.step({})
|
||||
continue
|
||||
agent_ids = list(obs_dict.keys())
|
||||
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||
with torch.no_grad():
|
||||
actions, _ = policy.sample(obs_tensor)
|
||||
actions = actions.cpu().numpy()
|
||||
action_dict = {aid: act for aid, act in zip(agent_ids, actions)}
|
||||
obs_dict, rewards, dones, _ = env.step(action_dict)
|
||||
episode_reward += sum(rewards.values())
|
||||
total_rewards.append(episode_reward)
|
||||
print(f" Eval Episode {i}: Total Reward {episode_reward:.2f}")
|
||||
avg_reward = float(np.mean(total_rewards))
|
||||
print(f" Average Evaluation Reward: {avg_reward:.2f}")
|
||||
return avg_reward
|
||||
except Exception as e:
|
||||
print(f"Evaluation failed: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return 0.0
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
def main(args):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f"Using device: {device}")
|
||||
|
||||
os.makedirs("logs/bc", exist_ok=True)
|
||||
log_dir = os.path.join("logs", "bc", datetime.now().strftime("%Y%m%d-%H%M%S"))
|
||||
writer = SummaryWriter(log_dir)
|
||||
print(f"TensorBoard logging to: {log_dir}")
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
|
||||
obs_data, act_data = load_expert_pkl(args.expert_data_path)
|
||||
obs_tensor = torch.FloatTensor(obs_data)
|
||||
act_tensor = torch.FloatTensor(act_data)
|
||||
dataset = TensorDataset(obs_tensor, act_tensor)
|
||||
train_size = int(0.8 * len(dataset))
|
||||
val_size = len(dataset) - train_size
|
||||
train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])
|
||||
train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True)
|
||||
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False)
|
||||
print(f"Dataset loaded. Train size: {len(train_dataset)}, Val size: {len(val_dataset)}")
|
||||
|
||||
state_dim = obs_data.shape[1]
|
||||
action_dim = act_data.shape[1]
|
||||
print(f"State Dim: {state_dim}, Action Dim: {action_dim}")
|
||||
|
||||
policy = StateIndependentPolicy(
|
||||
state_shape=(state_dim,),
|
||||
action_shape=(action_dim,),
|
||||
hidden_units=(256, 256),
|
||||
hidden_activation=torch.nn.Tanh(),
|
||||
).to(device)
|
||||
optimizer = Adam(policy.parameters(), lr=args.lr)
|
||||
scheduler = ExponentialLR(optimizer, gamma=0.99)
|
||||
|
||||
best_val_loss = float("inf")
|
||||
for epoch in range(args.epochs):
|
||||
avg_train_loss = train_bc_epoch(policy, train_loader, optimizer, device)
|
||||
scheduler.step()
|
||||
avg_val_loss = eval_bc_epoch(policy, val_loader, device)
|
||||
|
||||
print(f"Epoch {epoch+1}/{args.epochs} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}")
|
||||
writer.add_scalar("Loss/train", avg_train_loss, epoch)
|
||||
writer.add_scalar("Loss/val", avg_val_loss, epoch)
|
||||
writer.add_scalar("Learning_rate", scheduler.get_last_lr()[0], epoch)
|
||||
|
||||
if avg_val_loss < best_val_loss:
|
||||
best_val_loss = avg_val_loss
|
||||
torch.save(policy.state_dict(), os.path.join(args.save_dir, "policy_best.pt"))
|
||||
|
||||
if (epoch + 1) % args.eval_freq == 0:
|
||||
eval_reward = evaluate_policy(policy, args, device)
|
||||
writer.add_scalar("Reward/eval", eval_reward, epoch)
|
||||
|
||||
torch.save(policy.state_dict(), os.path.join(args.save_dir, "policy_final.pt"))
|
||||
writer.close()
|
||||
print("Training finished.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--expert_data_path", type=str, default="data/training_data", help="Path to expert data pickle or directory")
|
||||
parser.add_argument("--save_dir", type=str, default="models/bc", help="Directory to save models")
|
||||
parser.add_argument("--epochs", type=int, default=100)
|
||||
parser.add_argument("--batch_size", type=int, default=64)
|
||||
parser.add_argument("--lr", type=float, default=3e-4)
|
||||
parser.add_argument("--eval_freq", type=int, default=10)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
159
train_magail.py
159
train_magail.py
@@ -9,7 +9,8 @@ import argparse
|
||||
import signal
|
||||
import sys
|
||||
from torch.utils.data import DataLoader
|
||||
from dataset.magail_dataset import MAGAILExpertDataset
|
||||
from dataset.loader import MAGAILExpertDataset
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
|
||||
# --- Networks ---
|
||||
|
||||
@@ -79,18 +80,30 @@ class PPO:
|
||||
self.K_epochs = K_epochs
|
||||
self.mse_loss = nn.MSELoss()
|
||||
|
||||
def _log_prob_from_dist(self, dist, pre_tanh_action):
|
||||
# Tanh-squashed Gaussian log-prob with correction term.
|
||||
log_prob = dist.log_prob(pre_tanh_action)
|
||||
correction = torch.log(1 - torch.tanh(pre_tanh_action) ** 2 + 1e-6)
|
||||
return (log_prob - correction).sum(dim=-1)
|
||||
|
||||
def select_action(self, state):
|
||||
with torch.no_grad():
|
||||
state = torch.FloatTensor(state).cuda()
|
||||
dist = self.actor(state)
|
||||
action = dist.sample()
|
||||
action_logprob = dist.log_prob(action).sum(dim=-1)
|
||||
return action.cpu().numpy(), action_logprob.cpu().numpy()
|
||||
pre_tanh_action = dist.sample()
|
||||
action = torch.tanh(pre_tanh_action)
|
||||
action_logprob = self._log_prob_from_dist(dist, pre_tanh_action)
|
||||
return (
|
||||
action.cpu().numpy(),
|
||||
action_logprob.cpu().numpy(),
|
||||
pre_tanh_action.cpu().numpy()
|
||||
)
|
||||
|
||||
def update(self, memory):
|
||||
# Convert memory to tensors
|
||||
states = torch.FloatTensor(np.array(memory['states'])).cuda()
|
||||
actions = torch.FloatTensor(np.array(memory['actions'])).cuda()
|
||||
pre_tanh_actions = torch.FloatTensor(np.array(memory['pre_tanh_actions'])).cuda()
|
||||
logprobs = torch.FloatTensor(np.array(memory['logprobs'])).cuda()
|
||||
rewards = torch.FloatTensor(np.array(memory['rewards'])).cuda()
|
||||
next_states = torch.FloatTensor(np.array(memory['next_states'])).cuda()
|
||||
@@ -124,7 +137,7 @@ class PPO:
|
||||
for _ in range(self.K_epochs):
|
||||
# Evaluating old actions and values :
|
||||
dist = self.actor(states)
|
||||
action_logprobs = dist.log_prob(actions).sum(dim=-1)
|
||||
action_logprobs = self._log_prob_from_dist(dist, pre_tanh_actions)
|
||||
dist_entropy = dist.entropy().sum(dim=-1)
|
||||
state_values = self.critic(states).squeeze()
|
||||
|
||||
@@ -152,12 +165,7 @@ class PPO:
|
||||
# --- Training Loop ---
|
||||
|
||||
def train(args):
|
||||
# 1. Setup Environment (Dummy for now, usually you run simulation here)
|
||||
# But for MAGAIL we need to collect generated trajectories.
|
||||
# We need the Env class to be importable.
|
||||
from Env.scenario_env import MultiAgentScenarioEnv
|
||||
from Env.simple_idm_policy import ConstantVelocityPolicy # Just for init
|
||||
|
||||
# 1. Setup Environment (45-dim obs via BCScenarioEnv)
|
||||
# Config for Env
|
||||
env_config = {
|
||||
"data_directory": args.data_dir,
|
||||
@@ -200,12 +208,7 @@ def train(args):
|
||||
yield batch
|
||||
expert_iter = cycle(expert_loader)
|
||||
|
||||
# 4. Initialize Env
|
||||
from Env.expert_replay_env import ExpertReplayEnv # Using ReplayEnv for config, but we need ScenarioEnv for simulation?
|
||||
# Actually we need MultiAgentScenarioEnv for interactive training, not Replay.
|
||||
from Env.scenario_env import MultiAgentScenarioEnv
|
||||
from Env.simple_idm_policy import ConstantVelocityPolicy # Placeholder policy for init
|
||||
|
||||
# 4. Initialize Env (BCScenarioEnv provides 45-dim obs)
|
||||
# 2. Setup Models
|
||||
# Determine state dim from environment if possible, or use fixed
|
||||
# Expert data has 45 dim?
|
||||
@@ -220,48 +223,48 @@ def train(args):
|
||||
# We need to inject that same logic into the training env, OR
|
||||
# subclass MultiAgentScenarioEnv in the training script to override observation.
|
||||
|
||||
class MAGAILScenarioEnv(MultiAgentScenarioEnv):
|
||||
def _get_all_obs(self):
|
||||
# Same logic as ExpertReplayEnv to ensure compatibility
|
||||
obs_dict = {}
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
# 1. Ego State
|
||||
ego_state = [
|
||||
vehicle.position[0], vehicle.position[1],
|
||||
vehicle.velocity[0], vehicle.velocity[1],
|
||||
vehicle.heading_theta
|
||||
]
|
||||
|
||||
# 2. Neighbors
|
||||
candidates = []
|
||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
if other_id == agent_id:
|
||||
continue
|
||||
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||
if dist < 30.0:
|
||||
candidates.append((dist, other_vehicle))
|
||||
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
top_10 = candidates[:10]
|
||||
|
||||
neighbor_feats = []
|
||||
for _, neighbor in top_10:
|
||||
neighbor_feats.extend([
|
||||
neighbor.position[0] - vehicle.position[0],
|
||||
neighbor.position[1] - vehicle.position[1],
|
||||
neighbor.velocity[0],
|
||||
neighbor.velocity[1]
|
||||
])
|
||||
|
||||
missing = 10 - len(top_10)
|
||||
if missing > 0:
|
||||
neighbor_feats.extend([0.0] * (4 * missing))
|
||||
|
||||
obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||
obs_dict[agent_id] = obs
|
||||
return obs_dict
|
||||
|
||||
env = MAGAILScenarioEnv(config=env_config, agent2policy={}) # Pass empty dict if we control all externally
|
||||
# class MAGAILScenarioEnv(MultiAgentScenarioEnv):
|
||||
# def _get_all_obs(self):
|
||||
# # Same logic as ExpertReplayEnv to ensure compatibility
|
||||
# obs_dict = {}
|
||||
# for agent_id, vehicle in self.controlled_agents.items():
|
||||
# # 1. Ego State
|
||||
# ego_state = [
|
||||
# vehicle.position[0], vehicle.position[1],
|
||||
# vehicle.velocity[0], vehicle.velocity[1],
|
||||
# vehicle.heading_theta
|
||||
# ]
|
||||
#
|
||||
# # 2. Neighbors
|
||||
# candidates = []
|
||||
# for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
# if other_id == agent_id:
|
||||
# continue
|
||||
# dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||
# if dist < 30.0:
|
||||
# candidates.append((dist, other_vehicle))
|
||||
#
|
||||
# candidates.sort(key=lambda x: x[0])
|
||||
# top_10 = candidates[:10]
|
||||
#
|
||||
# neighbor_feats = []
|
||||
# for _, neighbor in top_10:
|
||||
# neighbor_feats.extend([
|
||||
# neighbor.position[0] - vehicle.position[0],
|
||||
# neighbor.position[1] - vehicle.position[1],
|
||||
# neighbor.velocity[0],
|
||||
# neighbor.velocity[1]
|
||||
# ])
|
||||
#
|
||||
# missing = 10 - len(top_10)
|
||||
# if missing > 0:
|
||||
# neighbor_feats.extend([0.0] * (4 * missing))
|
||||
#
|
||||
# obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||
# obs_dict[agent_id] = obs
|
||||
# return obs_dict
|
||||
|
||||
env = BCScenarioEnv(env_config, agent2policy={}) # 45-dim obs
|
||||
|
||||
print("Starting training...")
|
||||
|
||||
@@ -277,7 +280,15 @@ def train(args):
|
||||
|
||||
for i_episode in range(args.max_episodes):
|
||||
# --- 1. Collect Rollouts (Interaction) ---
|
||||
memory = {'states': [], 'actions': [], 'logprobs': [], 'rewards': [], 'next_states': [], 'dones': []}
|
||||
memory = {
|
||||
'states': [],
|
||||
'actions': [],
|
||||
'pre_tanh_actions': [],
|
||||
'logprobs': [],
|
||||
'rewards': [],
|
||||
'next_states': [],
|
||||
'dones': []
|
||||
}
|
||||
|
||||
# Prepare seed
|
||||
available_scenarios = env.config["num_scenarios"]
|
||||
@@ -338,7 +349,7 @@ def train(args):
|
||||
import gc
|
||||
gc.collect()
|
||||
|
||||
env = MAGAILScenarioEnv(config=env_config, agent2policy={})
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
obs_dict = env.reset(seed=seed)
|
||||
|
||||
episode_reward = 0
|
||||
@@ -349,6 +360,7 @@ def train(args):
|
||||
# Select actions for all agents
|
||||
actions = {}
|
||||
action_logprobs = {}
|
||||
pre_tanh_actions = {}
|
||||
|
||||
# obs_dict: {agent_id: obs}
|
||||
# MultiAgentScenarioEnv usually returns a dict {agent_id: obs}
|
||||
@@ -386,9 +398,10 @@ def train(args):
|
||||
obs_dict = new_obs_dict
|
||||
|
||||
for agent_id, obs in obs_dict.items():
|
||||
act, logprob = ppo_agent.select_action(obs) # Select action returns numpy
|
||||
act, logprob, pre_tanh = ppo_agent.select_action(obs) # Select action returns numpy
|
||||
actions[agent_id] = act.flatten() # (2,)
|
||||
action_logprobs[agent_id] = logprob # scalar
|
||||
pre_tanh_actions[agent_id] = pre_tanh.flatten()
|
||||
|
||||
# Step Env
|
||||
next_obs_dict, rewards, dones, infos = env.step(actions)
|
||||
@@ -398,6 +411,7 @@ def train(args):
|
||||
if agent_id in actions:
|
||||
memory['states'].append(obs)
|
||||
memory['actions'].append(actions[agent_id])
|
||||
memory['pre_tanh_actions'].append(pre_tanh_actions[agent_id])
|
||||
memory['logprobs'].append(action_logprobs[agent_id])
|
||||
|
||||
# Store standard environmental reward for logging (not used for update in GAIL)
|
||||
@@ -407,7 +421,7 @@ def train(args):
|
||||
# Next state
|
||||
if agent_id in next_obs_dict:
|
||||
memory['next_states'].append(next_obs_dict[agent_id])
|
||||
memory['dones'].append(False)
|
||||
memory['dones'].append(dones.get("__all__", False))
|
||||
else:
|
||||
# Agent finished/vanished
|
||||
# We need a dummy next state or handle done correctly
|
||||
@@ -460,6 +474,10 @@ def train(args):
|
||||
disc_loss = exp_loss + pol_loss
|
||||
disc_loss.backward()
|
||||
disc_optimizer.step()
|
||||
|
||||
with torch.no_grad():
|
||||
disc_acc_exp = (exp_preds > 0.5).float().mean().item()
|
||||
disc_acc_pol = (pol_preds < 0.5).float().mean().item()
|
||||
|
||||
# --- 3. Update Policy with GAIL Rewards ---
|
||||
# Reward = -log(1 - D(s, a))
|
||||
@@ -494,11 +512,18 @@ def train(args):
|
||||
writer.add_scalar('Loss/Discriminator', disc_loss.item(), i_episode)
|
||||
writer.add_scalar('Loss/Policy', ppo_loss, i_episode)
|
||||
writer.add_scalar('Reward/Mean_GAIL', np.mean(all_gail_rewards), i_episode)
|
||||
if batch_size > 0:
|
||||
writer.add_scalar('Acc/Disc_Expert', disc_acc_exp, i_episode)
|
||||
writer.add_scalar('Acc/Disc_Policy', disc_acc_pol, i_episode)
|
||||
if len(memory['actions']) > 0:
|
||||
action_arr = np.array(memory['actions'])
|
||||
action_clip_ratio = (np.abs(action_arr) > 0.98).mean()
|
||||
writer.add_scalar('Policy/ActionClipRatio', action_clip_ratio, i_episode)
|
||||
|
||||
print(f"Episode {i_episode}: Disc Loss {disc_loss.item():.4f} | PPO Loss {ppo_loss:.4f} | Mean Reward {np.mean(all_gail_rewards):.4f}")
|
||||
|
||||
if i_episode % 50 == 0:
|
||||
ppo_agent.save(os.path.join(args.log_dir, f"model_{i_episode}"))
|
||||
ppo_agent.save(os.path.join(args.save_dir, f"model_{i_episode}"))
|
||||
|
||||
env.close()
|
||||
if writer:
|
||||
@@ -511,11 +536,13 @@ if __name__ == '__main__':
|
||||
parser.add_argument("--batch_size", type=int, default=1024)
|
||||
parser.add_argument("--max_episodes", type=int, default=1000)
|
||||
parser.add_argument("--num_scenarios", type=int, default=100)
|
||||
parser.add_argument("--log_dir", type=str, default="runs/magail_exp")
|
||||
parser.add_argument("--log_dir", type=str, default="logs/magail", help="TensorBoard log directory")
|
||||
parser.add_argument("--save_dir", type=str, default="models/magail", help="Directory to save model checkpoints")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create log dir
|
||||
# Create log dir and save dir
|
||||
os.makedirs(args.log_dir, exist_ok=True)
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
|
||||
train(args)
|
||||
|
||||
Reference in New Issue
Block a user