Compare commits
11 Commits
train_not_
...
BC_SA
| Author | SHA1 | Date | |
|---|---|---|---|
| be35650533 | |||
| 8a75f0db0d | |||
| 0f9f080e77 | |||
| ceb6648a31 | |||
| 95cc78d940 | |||
| 03dee0205a | |||
| 21c046aef0 | |||
| 265b0eade1 | |||
| 4dbea5f0a6 | |||
| c94571ddaa | |||
| 62e638c4d2 |
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
|
||||
|
||||
0
Algorithm/__init__.py
Normal file
0
Algorithm/__init__.py
Normal file
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
|
||||
@@ -1,339 +0,0 @@
|
||||
# 调试功能使用指南
|
||||
|
||||
## 📋 概述
|
||||
|
||||
已为车道过滤和红绿灯检测功能添加了详细的调试输出,帮助您诊断和理解代码行为。
|
||||
|
||||
---
|
||||
|
||||
## 🎛️ 调试开关
|
||||
|
||||
### 1. 配置参数
|
||||
|
||||
在创建环境时,可以通过 `config` 参数启用调试模式:
|
||||
|
||||
```python
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
# ... 其他配置 ...
|
||||
|
||||
# 🔥 调试开关
|
||||
"debug_lane_filter": True, # 启用车道过滤调试
|
||||
"debug_traffic_light": True, # 启用红绿灯检测调试
|
||||
},
|
||||
agent2policy=your_policy
|
||||
)
|
||||
```
|
||||
|
||||
### 2. 默认值
|
||||
|
||||
两个调试开关默认都是 `False`(关闭),避免正常运行时产生大量日志。
|
||||
|
||||
---
|
||||
|
||||
## 📊 车道过滤调试 (`debug_lane_filter=True`)
|
||||
|
||||
### 输出内容
|
||||
|
||||
```
|
||||
📍 场景信息统计:
|
||||
- 总车道数: 123
|
||||
|
||||
🔍 开始车道过滤: 共 51 辆车待检测
|
||||
|
||||
车辆 1/51: ID=128
|
||||
🔍 检测位置 (-4.11, 46.76), 容差=3.0m
|
||||
✅ 在车道上 (车道184, 检查了32条)
|
||||
✅ 保留
|
||||
|
||||
车辆 7/51: ID=134
|
||||
🔍 检测位置 (-51.34, -3.77), 容差=3.0m
|
||||
❌ 不在任何车道上 (检查了123条车道)
|
||||
❌ 过滤 (原因: 不在车道上)
|
||||
|
||||
... (所有车辆)
|
||||
|
||||
📊 过滤结果: 保留 45 辆, 过滤 6 辆
|
||||
```
|
||||
|
||||
### 调试信息说明
|
||||
|
||||
| 信息 | 含义 |
|
||||
|------|------|
|
||||
| 📍 场景信息统计 | 场景的基本信息(车道数、红绿灯数) |
|
||||
| 🔍 开始车道过滤 | 开始过滤,显示待检测车辆总数 |
|
||||
| 🔍 检测位置 | 车辆的坐标和使用的容差值 |
|
||||
| ✅ 在车道上 | 找到了车辆所在的车道,显示车道ID和检查次数 |
|
||||
| ❌ 不在任何车道上 | 所有车道都检查完了,未找到匹配的车道 |
|
||||
| 📊 过滤结果 | 最终统计:保留多少辆,过滤多少辆 |
|
||||
|
||||
### 典型输出案例
|
||||
|
||||
**情况1:车辆在正常车道上**
|
||||
```
|
||||
车辆 1/51: ID=128
|
||||
🔍 检测位置 (-4.11, 46.76), 容差=3.0m
|
||||
✅ 在车道上 (车道184, 检查了32条)
|
||||
✅ 保留
|
||||
```
|
||||
→ 检查了32条车道后找到匹配的车道184
|
||||
|
||||
**情况2:车辆在草坪/停车场**
|
||||
```
|
||||
车辆 7/51: ID=134
|
||||
🔍 检测位置 (-51.34, -3.77), 容差=3.0m
|
||||
❌ 不在任何车道上 (检查了123条车道)
|
||||
❌ 过滤 (原因: 不在车道上)
|
||||
```
|
||||
→ 检查了所有123条车道都不匹配,该车辆被过滤
|
||||
|
||||
---
|
||||
|
||||
## 🚦 红绿灯检测调试 (`debug_traffic_light=True`)
|
||||
|
||||
### 输出内容
|
||||
|
||||
```
|
||||
📍 场景信息统计:
|
||||
- 总车道数: 123
|
||||
- 有红绿灯的车道数: 0
|
||||
⚠️ 场景中没有红绿灯!
|
||||
|
||||
🚦 检测车辆红绿灯 - 位置: (-4.1, 46.8)
|
||||
方法1-导航模块:
|
||||
current_lane = <metadrive.component.lane.straight_lane.StraightLane object>
|
||||
lane_index = 184
|
||||
has_traffic_light = False
|
||||
该车道没有红绿灯
|
||||
方法2-遍历车道: 开始遍历 123 条车道
|
||||
✓ 找到车辆所在车道: 184 (检查了32条)
|
||||
has_traffic_light = False
|
||||
该车道没有红绿灯
|
||||
结果: 返回 0 (无红绿灯/未知)
|
||||
```
|
||||
|
||||
### 调试信息说明
|
||||
|
||||
| 信息 | 含义 |
|
||||
|------|------|
|
||||
| 有红绿灯的车道数 | 统计场景中有多少个红绿灯 |
|
||||
| ⚠️ 场景中没有红绿灯 | 如果数量为0,会特别提示 |
|
||||
| 方法1-导航模块 | 尝试从导航系统获取 |
|
||||
| current_lane | 导航系统返回的当前车道对象 |
|
||||
| lane_index | 车道的唯一标识符 |
|
||||
| has_traffic_light | 该车道是否有红绿灯 |
|
||||
| status | 红绿灯的状态(GREEN/YELLOW/RED/None) |
|
||||
| 方法2-遍历车道 | 兜底方案,遍历所有车道查找 |
|
||||
| ✓ 找到车辆所在车道 | 遍历找到了匹配的车道 |
|
||||
|
||||
### 典型输出案例
|
||||
|
||||
**情况1:场景没有红绿灯**
|
||||
```
|
||||
📍 场景信息统计:
|
||||
- 有红绿灯的车道数: 0
|
||||
⚠️ 场景中没有红绿灯!
|
||||
|
||||
🚦 检测车辆红绿灯 - 位置: (-4.1, 46.8)
|
||||
方法1-导航模块:
|
||||
...
|
||||
has_traffic_light = False
|
||||
该车道没有红绿灯
|
||||
结果: 返回 0 (无红绿灯/未知)
|
||||
```
|
||||
→ 所有车辆都会返回0,这是正常的
|
||||
|
||||
**情况2:有红绿灯且状态正常**
|
||||
```
|
||||
🚦 检测车辆红绿灯 - 位置: (10.5, 20.3)
|
||||
方法1-导航模块:
|
||||
current_lane = <...>
|
||||
lane_index = 205
|
||||
has_traffic_light = True
|
||||
status = TRAFFIC_LIGHT_GREEN
|
||||
✅ 方法1成功: 绿灯
|
||||
```
|
||||
→ 方法1直接成功,返回1(绿灯)
|
||||
|
||||
**情况3:红绿灯状态为None**
|
||||
```
|
||||
🚦 检测车辆红绿灯 - 位置: (10.5, 20.3)
|
||||
方法1-导航模块:
|
||||
current_lane = <...>
|
||||
lane_index = 205
|
||||
has_traffic_light = True
|
||||
status = None
|
||||
⚠️ 方法1: 红绿灯状态为None
|
||||
```
|
||||
→ 有红绿灯,但状态异常,返回0
|
||||
|
||||
**情况4:导航失败,方法2兜底**
|
||||
```
|
||||
🚦 检测车辆红绿灯 - 位置: (15.2, 30.5)
|
||||
方法1-导航模块: 不可用 (hasattr=True, not_none=False)
|
||||
方法2-遍历车道: 开始遍历 123 条车道
|
||||
✓ 找到车辆所在车道: 210 (检查了45条)
|
||||
has_traffic_light = True
|
||||
status = TRAFFIC_LIGHT_RED
|
||||
✅ 方法2成功: 红灯
|
||||
```
|
||||
→ 方法1失败,方法2兜底成功,返回3(红灯)
|
||||
|
||||
---
|
||||
|
||||
## 🧪 测试方法
|
||||
|
||||
### 方式1:使用测试脚本
|
||||
|
||||
```bash
|
||||
# 标准测试(无详细调试)
|
||||
python Env/test_lane_filter.py
|
||||
|
||||
# 调试模式(详细输出)
|
||||
python Env/test_lane_filter.py --debug
|
||||
```
|
||||
|
||||
### 方式2:在代码中直接启用
|
||||
|
||||
```python
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": "...",
|
||||
"use_render": False,
|
||||
|
||||
# 启用调试
|
||||
"debug_lane_filter": True,
|
||||
"debug_traffic_light": True,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
obs = env.reset(0)
|
||||
# 调试信息会自动输出
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📝 调试输出控制
|
||||
|
||||
### 场景1:只想看车道过滤
|
||||
|
||||
```python
|
||||
config = {
|
||||
"debug_lane_filter": True,
|
||||
"debug_traffic_light": False, # 关闭红绿灯调试
|
||||
}
|
||||
```
|
||||
|
||||
### 场景2:只想看红绿灯检测
|
||||
|
||||
```python
|
||||
config = {
|
||||
"debug_lane_filter": False,
|
||||
"debug_traffic_light": True, # 只看红绿灯
|
||||
}
|
||||
```
|
||||
|
||||
### 场景3:生产环境(关闭所有调试)
|
||||
|
||||
```python
|
||||
config = {
|
||||
"debug_lane_filter": False,
|
||||
"debug_traffic_light": False,
|
||||
}
|
||||
# 或者直接不设置这两个参数,默认就是False
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 💡 常见问题诊断
|
||||
|
||||
### 问题1:所有红绿灯状态都是0
|
||||
|
||||
**检查调试输出:**
|
||||
```
|
||||
📍 场景信息统计:
|
||||
- 有红绿灯的车道数: 0
|
||||
⚠️ 场景中没有红绿灯!
|
||||
```
|
||||
|
||||
**结论:** 场景本身没有红绿灯,返回0是正常的
|
||||
|
||||
---
|
||||
|
||||
### 问题2:车辆被过滤但不应该过滤
|
||||
|
||||
**检查调试输出:**
|
||||
```
|
||||
车辆 X: ID=XXX
|
||||
🔍 检测位置 (x, y), 容差=3.0m
|
||||
❌ 不在任何车道上 (检查了123条车道)
|
||||
❌ 过滤 (原因: 不在车道上)
|
||||
```
|
||||
|
||||
**可能原因:**
|
||||
1. 车辆确实在非车道区域(草坪/停车场)
|
||||
2. 容差值太小,可以尝试增大 `lane_tolerance`
|
||||
3. 车道数据有问题
|
||||
|
||||
**解决方案:**
|
||||
```python
|
||||
config = {
|
||||
"lane_tolerance": 5.0, # 增大容差到5米
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 问题3:性能下降
|
||||
|
||||
启用调试模式会有大量输出,影响性能:
|
||||
|
||||
**解决方案:**
|
||||
- 只在开发/调试时启用
|
||||
- 生产环境关闭所有调试开关
|
||||
- 或者只测试少量车辆:
|
||||
```python
|
||||
config = {
|
||||
"max_controlled_vehicles": 5, # 只测试5辆车
|
||||
"debug_traffic_light": True,
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📌 最佳实践
|
||||
|
||||
1. **开发阶段**:启用调试,理解代码行为
|
||||
2. **调试问题**:根据需要选择性启用调试
|
||||
3. **性能测试**:关闭所有调试
|
||||
4. **生产运行**:永久关闭调试
|
||||
|
||||
---
|
||||
|
||||
## 🔧 调试输出示例
|
||||
|
||||
完整的调试运行示例:
|
||||
|
||||
```bash
|
||||
cd /home/huangfukk/MAGAIL4AutoDrive
|
||||
python Env/test_lane_filter.py --debug
|
||||
```
|
||||
|
||||
输出会包含:
|
||||
- 场景统计信息
|
||||
- 每辆车的详细检测过程
|
||||
- 最终的过滤/检测结果
|
||||
- 性能统计
|
||||
|
||||
---
|
||||
|
||||
## 📖 相关文档
|
||||
|
||||
- `README.md` - 项目总览和问题解决
|
||||
- `CHANGELOG.md` - 更新日志
|
||||
- `PERFORMANCE_OPTIMIZATION.md` - 性能优化指南
|
||||
|
||||
@@ -1,221 +0,0 @@
|
||||
# GPU加速指南
|
||||
|
||||
## 当前性能瓶颈分析
|
||||
|
||||
从测试结果看,即使关闭渲染,FPS仍然只有15-20左右,主要瓶颈是:
|
||||
|
||||
### 计算量分析(51辆车)
|
||||
```
|
||||
激光雷达计算:
|
||||
- 前向雷达:80束 × 51车 = 4,080次射线检测
|
||||
- 侧向雷达:10束 × 51车 = 510次射线检测
|
||||
- 车道线雷达:10束 × 51车 = 510次射线检测
|
||||
合计:5,100次射线检测/帧
|
||||
|
||||
红绿灯检测:
|
||||
- 遍历所有车道 × 51车 = 数千次几何计算
|
||||
```
|
||||
|
||||
**关键问题**:这些计算都是CPU单线程串行的,无法利用多核和GPU!
|
||||
|
||||
---
|
||||
|
||||
## GPU加速方案
|
||||
|
||||
### 方案1:优化激光雷达计算(已实现)✅
|
||||
|
||||
**优化内容:**
|
||||
1. 减少激光束数量:100束 → 52束(减少48%)
|
||||
2. 优化红绿灯检测:避免遍历所有车道
|
||||
3. 激光雷达缓存:每N帧才重新计算一次
|
||||
|
||||
**预期提升:** 2-4倍(30-60 FPS)
|
||||
|
||||
**使用方法:**
|
||||
```bash
|
||||
python Env/run_multiagent_env_fast.py
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 方案2:MetaDrive GPU渲染(有限支持)
|
||||
|
||||
**说明:**
|
||||
MetaDrive基于Panda3D引擎,理论上支持GPU渲染,但:
|
||||
- GPU主要用于**图形渲染**,不是物理计算
|
||||
- 激光雷达的射线检测仍在CPU上
|
||||
- GPU渲染主要加速可视化,不加速训练
|
||||
|
||||
**启用方法:**
|
||||
```python
|
||||
config = {
|
||||
"use_render": True,
|
||||
"render_mode": "onscreen", # 或 "offscreen"
|
||||
# Panda3D会自动尝试使用GPU
|
||||
}
|
||||
```
|
||||
|
||||
**限制:**
|
||||
- 需要显示器或虚拟显示(Xvfb)
|
||||
- WSL2环境需要配置X11转发
|
||||
- 对无渲染训练无帮助
|
||||
|
||||
---
|
||||
|
||||
### 方案3:使用GPU加速的物理引擎(推荐但需要迁移)
|
||||
|
||||
**选项A:Isaac Gym (NVIDIA)**
|
||||
- 完全在GPU上运行物理模拟和渲染
|
||||
- 可同时模拟数千个环境
|
||||
- **缺点**:需要完全重写环境代码,迁移成本高
|
||||
|
||||
**选项B:IsaacSim/Omniverse**
|
||||
- NVIDIA的高级仿真平台
|
||||
- 支持GPU加速的激光雷达
|
||||
- **缺点**:学习曲线陡峭,环境配置复杂
|
||||
|
||||
**选项C:Brax (Google)**
|
||||
- JAX驱动,完全在GPU/TPU上运行
|
||||
- **缺点**:功能有限,不支持复杂场景
|
||||
|
||||
---
|
||||
|
||||
### 方案4:策略网络GPU加速(推荐)✅
|
||||
|
||||
虽然环境仿真在CPU,但可以让**策略网络在GPU上运行**:
|
||||
|
||||
```python
|
||||
import torch
|
||||
|
||||
# 创建GPU上的策略模型
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
policy = PolicyNetwork().to(device)
|
||||
|
||||
# 批量处理观测
|
||||
obs_batch = torch.tensor(obs_list).to(device)
|
||||
with torch.no_grad():
|
||||
actions = policy(obs_batch)
|
||||
actions = actions.cpu().numpy()
|
||||
```
|
||||
|
||||
**优势:**
|
||||
- 51辆车的推理可以并行
|
||||
- 如果使用RL训练,GPU加速训练过程
|
||||
- 不需要修改环境代码
|
||||
|
||||
---
|
||||
|
||||
### 方案5:多进程并行(最实用)✅
|
||||
|
||||
既然单个环境受限于CPU单线程,可以**并行运行多个环境**:
|
||||
|
||||
```python
|
||||
from multiprocessing import Pool
|
||||
import os
|
||||
|
||||
def run_single_env(seed):
|
||||
"""运行单个环境实例"""
|
||||
env = MultiAgentScenarioEnv(config=...)
|
||||
obs = env.reset(seed)
|
||||
|
||||
for step in range(1000):
|
||||
actions = {...}
|
||||
obs, rewards, dones, infos = env.step(actions)
|
||||
if dones["__all__"]:
|
||||
break
|
||||
|
||||
env.close()
|
||||
return results
|
||||
|
||||
# 使用进程池并行运行
|
||||
if __name__ == "__main__":
|
||||
num_processes = os.cpu_count() # 12600KF有10核20线程
|
||||
seeds = list(range(num_processes))
|
||||
|
||||
with Pool(processes=num_processes) as pool:
|
||||
results = pool.map(run_single_env, seeds)
|
||||
```
|
||||
|
||||
**预期提升:** 接近线性(10核 ≈ 10倍吞吐量)
|
||||
|
||||
**CPU利用率:** 可达80-100%
|
||||
|
||||
---
|
||||
|
||||
## 推荐的完整优化方案
|
||||
|
||||
### 1. 立即可用(已实现)
|
||||
```bash
|
||||
# 使用优化版本,激光束减少+缓存
|
||||
python Env/run_multiagent_env_fast.py
|
||||
```
|
||||
**预期:** 30-60 FPS(2-4倍提升)
|
||||
|
||||
### 2. 短期优化(1-2小时)
|
||||
- 实现多进程并行
|
||||
- 策略网络迁移到GPU
|
||||
|
||||
**预期:** 300-600 FPS(总吞吐量)
|
||||
|
||||
### 3. 中期优化(1-2天)
|
||||
- 使用NumPy矢量化批量处理观测
|
||||
- 优化Python代码热点(用Cython/Numba)
|
||||
|
||||
**预期:** 额外20-30%提升
|
||||
|
||||
### 4. 长期方案(1-2周)
|
||||
- 迁移到Isaac Gym等GPU加速仿真器
|
||||
- 或使用分布式训练框架(Ray/RLlib)
|
||||
|
||||
**预期:** 10-100倍提升
|
||||
|
||||
---
|
||||
|
||||
## 为什么MetaDrive无法直接使用GPU?
|
||||
|
||||
### 架构限制:
|
||||
1. **物理引擎**:使用Bullet/Panda3D的CPU物理引擎
|
||||
2. **射线检测**:串行CPU计算,无法并行
|
||||
3. **Python GIL**:全局解释器锁限制多线程
|
||||
4. **设计目标**:MetaDrive设计时主要考虑灵活性而非极致性能
|
||||
|
||||
### GPU在仿真中的作用:
|
||||
- ✅ **图形渲染**:绘制画面(但我们训练时不需要)
|
||||
- ✅ **神经网络推理/训练**:策略模型计算
|
||||
- ❌ **物理计算**:MetaDrive的物理引擎在CPU
|
||||
- ❌ **传感器模拟**:激光雷达等在CPU
|
||||
|
||||
---
|
||||
|
||||
## 检查GPU是否可用
|
||||
|
||||
```bash
|
||||
# 检查NVIDIA GPU
|
||||
nvidia-smi
|
||||
|
||||
# 检查PyTorch GPU支持
|
||||
python -c "import torch; print(f'CUDA available: {torch.cuda.is_available()}')"
|
||||
|
||||
# 检查MetaDrive渲染设备
|
||||
python -c "from panda3d.core import GraphicsPipeSelection; print(GraphicsPipeSelection.get_global_ptr().get_default_pipe())"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 总结
|
||||
|
||||
| 方案 | 实现难度 | 性能提升 | GPU使用 | 推荐度 |
|
||||
|------|----------|----------|---------|--------|
|
||||
| 减少激光束 | ⭐ | 2-4x | ❌ | ⭐⭐⭐⭐⭐ |
|
||||
| 激光雷达缓存 | ⭐ | 1.5-3x | ❌ | ⭐⭐⭐⭐⭐ |
|
||||
| 多进程并行 | ⭐⭐ | 5-10x | ❌ | ⭐⭐⭐⭐⭐ |
|
||||
| 策略GPU加速 | ⭐⭐ | 2-5x | ✅ | ⭐⭐⭐⭐ |
|
||||
| GPU渲染 | ⭐⭐⭐ | 1.2x | ✅ | ⭐⭐ |
|
||||
| 迁移Isaac Gym | ⭐⭐⭐⭐⭐ | 10-100x | ✅ | ⭐⭐⭐ |
|
||||
|
||||
**结论:**
|
||||
1. 先用已实现的优化(减少激光束+缓存)
|
||||
2. 再实现多进程并行
|
||||
3. 策略网络用GPU训练
|
||||
4. 如果还不够,考虑迁移到GPU仿真器
|
||||
|
||||
@@ -1,413 +0,0 @@
|
||||
# 日志记录功能使用指南
|
||||
|
||||
## 📋 概述
|
||||
|
||||
为所有运行脚本添加了日志记录功能,可以将终端输出同时保存到文本文件,方便后续分析和问题排查。
|
||||
|
||||
---
|
||||
|
||||
## 🎯 功能特点
|
||||
|
||||
1. **双向输出**:同时输出到终端和文件,不影响实时查看
|
||||
2. **自动管理**:使用上下文管理器,自动处理文件开启/关闭
|
||||
3. **灵活配置**:支持自定义文件名和日志目录
|
||||
4. **时间戳命名**:默认使用时间戳生成唯一文件名
|
||||
5. **无缝集成**:只需添加命令行参数,无需修改代码
|
||||
|
||||
---
|
||||
|
||||
## 🚀 快速使用
|
||||
|
||||
### 1. 基础用法
|
||||
|
||||
```bash
|
||||
# 不启用日志(默认)
|
||||
python Env/run_multiagent_env.py
|
||||
|
||||
# 启用日志记录
|
||||
python Env/run_multiagent_env.py --log
|
||||
|
||||
# 或使用短选项
|
||||
python Env/run_multiagent_env.py -l
|
||||
```
|
||||
|
||||
### 2. 自定义文件名
|
||||
|
||||
```bash
|
||||
# 使用自定义日志文件名
|
||||
python Env/run_multiagent_env.py --log --log-file=my_test.log
|
||||
|
||||
# 测试脚本也支持
|
||||
python Env/test_lane_filter.py --log --log-file=test_results.log
|
||||
```
|
||||
|
||||
### 3. 组合使用调试和日志
|
||||
|
||||
```bash
|
||||
# 测试脚本:调试模式 + 日志记录
|
||||
python Env/test_lane_filter.py --debug --log
|
||||
|
||||
# 会生成类似:test_debug_20251021_123456.log
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📁 日志文件位置
|
||||
|
||||
默认日志目录:`Env/logs/`
|
||||
|
||||
### 文件命名规则
|
||||
|
||||
| 脚本 | 默认文件名格式 | 示例 |
|
||||
|------|---------------|------|
|
||||
| `run_multiagent_env.py` | `run_YYYYMMDD_HHMMSS.log` | `run_20251021_143022.log` |
|
||||
| `run_multiagent_env_fast.py` | `run_fast.log` | `run_fast.log` |
|
||||
| `test_lane_filter.py` | `test_{mode}_YYYYMMDD_HHMMSS.log` | `test_debug_20251021_143500.log` |
|
||||
|
||||
**说明**:
|
||||
- `YYYYMMDD_HHMMSS` 是时间戳(年月日_时分秒)
|
||||
- `{mode}` 是测试模式(`standard` 或 `debug`)
|
||||
|
||||
---
|
||||
|
||||
## 📝 所有支持的脚本
|
||||
|
||||
### 1. run_multiagent_env.py(标准运行脚本)
|
||||
|
||||
```bash
|
||||
# 不启用日志
|
||||
python Env/run_multiagent_env.py
|
||||
|
||||
# 启用日志(自动生成时间戳文件名)
|
||||
python Env/run_multiagent_env.py --log
|
||||
|
||||
# 自定义文件名
|
||||
python Env/run_multiagent_env.py --log --log-file=run_test1.log
|
||||
```
|
||||
|
||||
**日志位置**:`Env/logs/run_YYYYMMDD_HHMMSS.log`
|
||||
|
||||
---
|
||||
|
||||
### 2. run_multiagent_env_fast.py(高性能版本)
|
||||
|
||||
```bash
|
||||
# 启用日志
|
||||
python Env/run_multiagent_env_fast.py --log
|
||||
|
||||
# 自定义文件名
|
||||
python Env/run_multiagent_env_fast.py --log --log-file=fast_test.log
|
||||
```
|
||||
|
||||
**日志位置**:`Env/logs/run_fast.log`(默认)
|
||||
|
||||
---
|
||||
|
||||
### 3. test_lane_filter.py(测试脚本)
|
||||
|
||||
```bash
|
||||
# 标准测试 + 日志
|
||||
python Env/test_lane_filter.py --log
|
||||
|
||||
# 调试测试 + 日志
|
||||
python Env/test_lane_filter.py --debug --log
|
||||
|
||||
# 自定义文件名
|
||||
python Env/test_lane_filter.py --log --log-file=my_test.log
|
||||
|
||||
# 组合使用
|
||||
python Env/test_lane_filter.py --debug --log --log-file=debug_run.log
|
||||
```
|
||||
|
||||
**日志位置**:
|
||||
- 标准模式:`Env/logs/test_standard_YYYYMMDD_HHMMSS.log`
|
||||
- 调试模式:`Env/logs/test_debug_YYYYMMDD_HHMMSS.log`
|
||||
|
||||
---
|
||||
|
||||
## 💻 编程接口
|
||||
|
||||
如果您想在代码中直接使用日志功能:
|
||||
|
||||
```python
|
||||
from logger_utils import setup_logger
|
||||
|
||||
# 方式1:使用上下文管理器(推荐)
|
||||
with setup_logger(log_file="my_log.log", log_dir="logs"):
|
||||
print("这条消息会同时输出到终端和文件")
|
||||
# 运行您的代码
|
||||
# ...
|
||||
|
||||
# 方式2:手动管理
|
||||
from logger_utils import LoggerContext
|
||||
|
||||
logger = LoggerContext(log_file="custom.log", log_dir="output")
|
||||
logger.__enter__() # 开启日志
|
||||
print("输出消息")
|
||||
logger.__exit__(None, None, None) # 关闭日志
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📊 日志内容示例
|
||||
|
||||
### 标准运行
|
||||
|
||||
```
|
||||
📝 日志记录已启用
|
||||
📁 日志文件: Env/logs/run_20251021_143022.log
|
||||
------------------------------------------------------------
|
||||
💡 提示: 使用 --log 或 -l 参数启用日志记录
|
||||
示例: python run_multiagent_env.py --log
|
||||
自定义文件名: python run_multiagent_env.py --log --log-file=my_run.log
|
||||
------------------------------------------------------------
|
||||
[INFO] Environment: MultiAgentScenarioEnv
|
||||
[INFO] MetaDrive version: 0.4.3
|
||||
...
|
||||
------------------------------------------------------------
|
||||
✅ 日志已保存到: Env/logs/run_20251021_143022.log
|
||||
```
|
||||
|
||||
### 调试模式
|
||||
|
||||
```
|
||||
📝 日志记录已启用
|
||||
📁 日志文件: Env/logs/test_debug_20251021_143500.log
|
||||
------------------------------------------------------------
|
||||
🐛 调试模式启用
|
||||
============================================================
|
||||
|
||||
📍 场景信息统计:
|
||||
- 总车道数: 123
|
||||
- 有红绿灯的车道数: 0
|
||||
⚠️ 场景中没有红绿灯!
|
||||
|
||||
🔍 开始车道过滤: 共 51 辆车待检测
|
||||
...
|
||||
------------------------------------------------------------
|
||||
✅ 日志已保存到: Env/logs/test_debug_20251021_143500.log
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🔧 高级配置
|
||||
|
||||
### 自定义日志目录
|
||||
|
||||
```python
|
||||
from logger_utils import setup_logger
|
||||
|
||||
# 指定不同的日志目录
|
||||
with setup_logger(log_file="test.log", log_dir="my_logs"):
|
||||
print("日志会保存到 my_logs/test.log")
|
||||
```
|
||||
|
||||
### 追加模式
|
||||
|
||||
```python
|
||||
from logger_utils import setup_logger
|
||||
|
||||
# 追加到现有文件(而不是覆盖)
|
||||
with setup_logger(log_file="test.log", mode='a'): # mode='a' 表示追加
|
||||
print("这条消息会追加到文件末尾")
|
||||
```
|
||||
|
||||
### 只重定向特定输出
|
||||
|
||||
```python
|
||||
from logger_utils import LoggerContext
|
||||
|
||||
# 只重定向stdout,不重定向stderr
|
||||
logger = LoggerContext(
|
||||
log_file="test.log",
|
||||
redirect_stdout=True, # 重定向标准输出
|
||||
redirect_stderr=False # 不重定向错误输出
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📋 命令行参数总结
|
||||
|
||||
| 参数 | 短选项 | 说明 | 示例 |
|
||||
|------|--------|------|------|
|
||||
| `--log` | `-l` | 启用日志记录 | `--log` |
|
||||
| `--log-file=NAME` | 无 | 指定日志文件名 | `--log-file=test.log` |
|
||||
| `--debug` | `-d` | 启用调试模式(test_lane_filter.py) | `--debug` |
|
||||
|
||||
### 参数组合
|
||||
|
||||
```bash
|
||||
# 示例1:标准模式 + 日志
|
||||
python Env/test_lane_filter.py --log
|
||||
|
||||
# 示例2:调试模式 + 日志
|
||||
python Env/test_lane_filter.py --debug --log
|
||||
|
||||
# 示例3:调试 + 自定义文件名
|
||||
python Env/test_lane_filter.py -d --log --log-file=my_debug.log
|
||||
|
||||
# 示例4:所有参数
|
||||
python Env/test_lane_filter.py --debug --log --log-file=full_test.log
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🛠️ 常见问题
|
||||
|
||||
### Q1: 日志文件在哪里?
|
||||
|
||||
**A**: 默认在 `Env/logs/` 目录下。如果目录不存在,会自动创建。
|
||||
|
||||
```bash
|
||||
# 查看所有日志文件
|
||||
ls -lh Env/logs/
|
||||
|
||||
# 查看最新的日志
|
||||
ls -lt Env/logs/ | head -5
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Q2: 如何查看日志内容?
|
||||
|
||||
**A**: 使用任何文本编辑器或命令行工具:
|
||||
|
||||
```bash
|
||||
# 方式1:使用cat
|
||||
cat Env/logs/run_20251021_143022.log
|
||||
|
||||
# 方式2:使用less(可翻页)
|
||||
less Env/logs/run_20251021_143022.log
|
||||
|
||||
# 方式3:查看末尾内容
|
||||
tail -n 50 Env/logs/run_20251021_143022.log
|
||||
|
||||
# 方式4:实时监控(适合长时间运行)
|
||||
tail -f Env/logs/run_20251021_143022.log
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Q3: 日志文件太多怎么办?
|
||||
|
||||
**A**: 可以定期清理旧日志:
|
||||
|
||||
```bash
|
||||
# 删除7天前的日志
|
||||
find Env/logs/ -name "*.log" -mtime +7 -delete
|
||||
|
||||
# 只保留最新的10个日志
|
||||
cd Env/logs && ls -t *.log | tail -n +11 | xargs rm -f
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Q4: 日志会影响性能吗?
|
||||
|
||||
**A**: 影响很小,因为:
|
||||
1. 文件I/O是异步的
|
||||
2. 使用了缓冲区
|
||||
3. 立即刷新确保数据不丢失
|
||||
|
||||
如果追求极致性能,建议训练时不启用日志,只在需要分析时启用。
|
||||
|
||||
---
|
||||
|
||||
### Q5: 可以同时记录多个脚本的日志吗?
|
||||
|
||||
**A**: 可以,每个脚本使用不同的日志文件:
|
||||
|
||||
```bash
|
||||
# 终端1
|
||||
python Env/run_multiagent_env.py --log --log-file=script1.log
|
||||
|
||||
# 终端2(同时运行)
|
||||
python Env/test_lane_filter.py --log --log-file=script2.log
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 💡 最佳实践
|
||||
|
||||
### 1. 开发阶段
|
||||
|
||||
```bash
|
||||
# 使用调试模式 + 日志,方便排查问题
|
||||
python Env/test_lane_filter.py --debug --log
|
||||
```
|
||||
|
||||
### 2. 长时间运行
|
||||
|
||||
```bash
|
||||
# 启用日志,避免输出丢失
|
||||
nohup python Env/run_multiagent_env.py --log > /dev/null 2>&1 &
|
||||
|
||||
# 查看实时输出
|
||||
tail -f Env/logs/run_*.log
|
||||
```
|
||||
|
||||
### 3. 批量实验
|
||||
|
||||
```bash
|
||||
# 为每次实验使用不同的日志文件
|
||||
for i in {1..5}; do
|
||||
python Env/run_multiagent_env.py --log --log-file=exp_${i}.log
|
||||
done
|
||||
```
|
||||
|
||||
### 4. 性能测试
|
||||
|
||||
```bash
|
||||
# 不启用日志,获得最佳性能
|
||||
python Env/run_multiagent_env_fast.py
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📖 相关文档
|
||||
|
||||
- `README.md` - 项目总览
|
||||
- `DEBUG_GUIDE.md` - 调试功能使用指南
|
||||
- `CHANGELOG.md` - 更新日志
|
||||
|
||||
---
|
||||
|
||||
## 🔍 技术细节
|
||||
|
||||
### 实现原理
|
||||
|
||||
1. **TeeLogger类**:实现同时写入终端和文件
|
||||
2. **上下文管理器**:自动管理资源(文件打开/关闭)
|
||||
3. **sys.stdout重定向**:拦截所有print输出
|
||||
4. **即时刷新**:每次写入后立即刷新,确保数据不丢失
|
||||
|
||||
### 源代码
|
||||
|
||||
详见 `Env/logger_utils.py`
|
||||
|
||||
```python
|
||||
# 简化示例
|
||||
class TeeLogger:
|
||||
def write(self, message):
|
||||
self.terminal.write(message) # 输出到终端
|
||||
self.log_file.write(message) # 写入文件
|
||||
self.log_file.flush() # 立即刷新
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## ✅ 总结
|
||||
|
||||
- ✅ 简单易用:只需添加 `--log` 参数
|
||||
- ✅ 不影响输出:终端仍可实时查看
|
||||
- ✅ 自动管理:文件自动开启/关闭
|
||||
- ✅ 灵活配置:支持自定义文件名和目录
|
||||
- ✅ 完整记录:包含所有调试信息
|
||||
|
||||
立即开始使用:
|
||||
|
||||
```bash
|
||||
python Env/test_lane_filter.py --debug --log
|
||||
```
|
||||
|
||||
@@ -1,131 +0,0 @@
|
||||
# MetaDrive 性能优化指南
|
||||
|
||||
## 为什么帧率只有15FPS且CPU利用率不高?
|
||||
|
||||
### 主要原因:
|
||||
|
||||
1. **渲染瓶颈(最主要)**
|
||||
- `use_render: True` + 每帧调用 `env.render()` 会严重限制帧率
|
||||
- MetaDrive 使用 Panda3D 渲染引擎,渲染是**同步阻塞**的
|
||||
- 即使CPU有余力,也要等待渲染完成才能继续下一步
|
||||
- 这就是为什么CPU利用率低但帧率也低的原因
|
||||
|
||||
2. **激光雷达计算开销**
|
||||
- 每帧对每辆车进行3次激光雷达扫描(100个激光束)
|
||||
- 需要进行物理射线检测,计算量较大
|
||||
|
||||
3. **物理引擎同步**
|
||||
- 默认物理步长很小(0.02s),需要频繁计算
|
||||
|
||||
4. **Python GIL限制**
|
||||
- Python全局解释器锁限制了多核并行
|
||||
- 即使是多核CPU,Python单线程性能才是瓶颈
|
||||
|
||||
## 性能优化方案
|
||||
|
||||
### 方案1:关闭渲染(推荐用于训练)
|
||||
**预期提升:10-20倍(150-300+ FPS)**
|
||||
|
||||
```python
|
||||
config = {
|
||||
"use_render": False, # 关闭渲染
|
||||
"render_pipeline": False,
|
||||
"image_observation": False,
|
||||
"interface_panel": [],
|
||||
"manual_control": False,
|
||||
}
|
||||
```
|
||||
|
||||
### 方案2:降低物理计算频率
|
||||
**预期提升:2-3倍**
|
||||
|
||||
```python
|
||||
config = {
|
||||
"physics_world_step_size": 0.05, # 默认0.02,增大步长
|
||||
"decision_repeat": 5, # 每5个物理步执行一次决策
|
||||
}
|
||||
```
|
||||
|
||||
### 方案3:优化激光雷达
|
||||
**预期提升:1.5-2倍**
|
||||
|
||||
修改 `scenario_env.py` 中的 `_get_all_obs()` 函数:
|
||||
|
||||
```python
|
||||
# 减少激光束数量
|
||||
lidar = self.engine.get_sensor("lidar").perceive(
|
||||
num_lasers=40, # 从80减到40
|
||||
distance=30,
|
||||
base_vehicle=vehicle,
|
||||
physics_world=self.engine.physics_world.dynamic_world
|
||||
)
|
||||
|
||||
# 或者降低扫描频率(每N步才扫描一次)
|
||||
if self.round % 5 == 0:
|
||||
lidar = self.engine.get_sensor("lidar").perceive(...)
|
||||
else:
|
||||
lidar = self.last_lidar[agent_id] # 使用缓存
|
||||
```
|
||||
|
||||
### 方案4:间歇性渲染
|
||||
**适用场景:既需要可视化又想提升性能**
|
||||
|
||||
```python
|
||||
# 每10步渲染一次,而不是每步都渲染
|
||||
if step % 10 == 0:
|
||||
env.render(mode="topdown")
|
||||
```
|
||||
|
||||
### 方案5:使用多进程并行(高级)
|
||||
**预期提升:接近线性(取决于进程数)**
|
||||
|
||||
```python
|
||||
from multiprocessing import Pool
|
||||
|
||||
def run_env(seed):
|
||||
env = MultiAgentScenarioEnv(config=...)
|
||||
# 运行仿真
|
||||
return results
|
||||
|
||||
# 使用进程池并行运行多个环境
|
||||
with Pool(processes=8) as pool:
|
||||
results = pool.map(run_env, range(8))
|
||||
```
|
||||
|
||||
## 文件说明
|
||||
|
||||
- `run_multiagent_env.py` - **标准版本**(无渲染,基础优化)
|
||||
- `run_multiagent_env_fast.py` - **极速版本**(激光雷达优化+缓存)⭐推荐
|
||||
- `run_multiagent_env_parallel.py` - **并行版本**(多进程,最高吞吐量)⭐⭐推荐
|
||||
- `run_multiagent_env_visual.py` - **可视化版本**(有渲染,适合调试)
|
||||
|
||||
## 性能对比
|
||||
|
||||
| 配置 | 单环境FPS | 总吞吐量 | CPU利用率 | 文件 | 适用场景 |
|
||||
|------|-----------|----------|-----------|------|----------|
|
||||
| 原始配置(有渲染) | 15-20 | 15-20 | 15-20% | visual | 实时可视化调试 |
|
||||
| 关闭渲染 | 20-25 | 20-25 | 20-30% | 标准版 | 基础训练 |
|
||||
| 激光雷达优化+缓存 | 30-60 | 30-60 | 30-50% | fast | 快速训练⭐ |
|
||||
| 多进程并行(10核) | 30-60 | 300-600 | 90-100% | parallel | 大规模训练⭐⭐ |
|
||||
|
||||
**说明:**
|
||||
- **单环境FPS**:单个环境实例的帧率
|
||||
- **总吞吐量**:所有进程合计的 steps/second
|
||||
- 12600KF(10核20线程)推荐使用并行版本
|
||||
|
||||
## 建议
|
||||
|
||||
1. **训练时**:使用高性能版本(关闭渲染)
|
||||
2. **调试时**:使用可视化版本,或间歇性渲染
|
||||
3. **大规模实验**:使用多进程并行
|
||||
4. **如果需要GPU加速**:考虑使用GPU渲染或将策略网络部署到GPU上
|
||||
|
||||
## 为什么CPU利用率低?
|
||||
|
||||
- **渲染阻塞**:CPU在等待渲染完成
|
||||
- **Python GIL**:限制了多核利用
|
||||
- **I/O等待**:可能在等待磁盘读取数据
|
||||
- **单线程瓶颈**:MetaDrive主循环是单线程的
|
||||
|
||||
解决方法:关闭渲染 + 多进程并行
|
||||
|
||||
@@ -1,241 +0,0 @@
|
||||
# 快速使用指南
|
||||
|
||||
## 🚀 已实现的性能优化
|
||||
|
||||
根据您的测试结果,原始版本FPS只有15左右,现已进行了全面优化。
|
||||
|
||||
---
|
||||
|
||||
## 📊 性能瓶颈分析
|
||||
|
||||
您的CPU是12600KF(10核20线程),但利用率不到20%,原因是:
|
||||
|
||||
1. **激光雷达计算瓶颈**:51辆车 × 100个激光束 = 每帧5100次射线检测
|
||||
2. **红绿灯检测低效**:遍历所有车道进行几何计算
|
||||
3. **Python GIL限制**:单线程执行,无法利用多核
|
||||
4. **计算串行化**:所有车辆依次处理,没有并行
|
||||
|
||||
---
|
||||
|
||||
## 🎯 推荐使用方案
|
||||
|
||||
### 方案1:极速单环境(推荐新手)⭐
|
||||
```bash
|
||||
python Env/run_multiagent_env_fast.py
|
||||
```
|
||||
|
||||
**优化内容:**
|
||||
- ✅ 激光束:100束 → 52束(减少48%计算量)
|
||||
- ✅ 激光雷达缓存:每3帧才重新计算
|
||||
- ✅ 红绿灯检测优化:避免遍历所有车道
|
||||
- ✅ 关闭所有渲染和调试
|
||||
|
||||
**预期性能:** 30-60 FPS(2-4倍提升)
|
||||
|
||||
---
|
||||
|
||||
### 方案2:多进程并行(推荐训练)⭐⭐
|
||||
```bash
|
||||
python Env/run_multiagent_env_parallel.py
|
||||
```
|
||||
|
||||
**优化内容:**
|
||||
- ✅ 同时运行10个独立环境(充分利用10核CPU)
|
||||
- ✅ 每个环境应用所有单环境优化
|
||||
- ✅ CPU利用率可达90-100%
|
||||
|
||||
**预期性能:** 300-600 steps/s(20-40倍总吞吐量)
|
||||
|
||||
---
|
||||
|
||||
### 方案3:可视化调试
|
||||
```bash
|
||||
python Env/run_multiagent_env_visual.py
|
||||
```
|
||||
|
||||
**说明:** 保留渲染功能,FPS约15,仅用于调试
|
||||
|
||||
---
|
||||
|
||||
## 🔧 关于GPU加速
|
||||
|
||||
### GPU能否加速MetaDrive?
|
||||
|
||||
**简短回答:有限支持,主要瓶颈不在GPU**
|
||||
|
||||
**详细说明:**
|
||||
|
||||
1. **物理计算(主要瓶颈)** ❌ 不支持GPU
|
||||
- MetaDrive使用Bullet物理引擎,只在CPU运行
|
||||
- 激光雷达射线检测也在CPU
|
||||
- 这是FPS低的主要原因
|
||||
|
||||
2. **图形渲染** ✅ 支持GPU
|
||||
- Panda3D会自动使用GPU渲染
|
||||
- 但我们训练时关闭了渲染,所以GPU无用武之地
|
||||
|
||||
3. **策略网络** ✅ 支持GPU
|
||||
- 可以把Policy模型放到GPU上
|
||||
- 但环境本身仍在CPU
|
||||
|
||||
### GPU渲染配置(可选)
|
||||
```python
|
||||
config = {
|
||||
"use_render": True,
|
||||
# GPU会自动用于渲染
|
||||
}
|
||||
```
|
||||
|
||||
### 策略网络GPU加速(推荐)
|
||||
```python
|
||||
import torch
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
policy_model = PolicyNet().to(device)
|
||||
|
||||
# 批量推理
|
||||
obs_tensor = torch.tensor(obs_list).to(device)
|
||||
actions = policy_model(obs_tensor)
|
||||
```
|
||||
|
||||
**详细说明请看:** `GPU_ACCELERATION.md`
|
||||
|
||||
---
|
||||
|
||||
## 📈 性能对比
|
||||
|
||||
| 版本 | FPS | CPU利用率 | 改进 |
|
||||
|------|-----|-----------|------|
|
||||
| 原始版本 | 15 | 20% | - |
|
||||
| 极速版本 | 30-60 | 30-50% | 2-4x |
|
||||
| 并行版本 | 30-60/env | 90-100% | 总吞吐20-40x |
|
||||
|
||||
---
|
||||
|
||||
## 💡 使用建议
|
||||
|
||||
### 场景1:快速测试环境
|
||||
```bash
|
||||
python Env/run_multiagent_env_fast.py
|
||||
```
|
||||
单环境,快速验证功能
|
||||
|
||||
### 场景2:大规模数据收集
|
||||
```bash
|
||||
python Env/run_multiagent_env_parallel.py
|
||||
```
|
||||
多进程,最大化数据收集速度
|
||||
|
||||
### 场景3:RL训练
|
||||
```bash
|
||||
# 推荐使用Ray RLlib等框架,它们内置了并行环境管理
|
||||
# 或者修改parallel版本,保存经验到replay buffer
|
||||
```
|
||||
|
||||
### 场景4:调试/可视化
|
||||
```bash
|
||||
python Env/run_multiagent_env_visual.py
|
||||
```
|
||||
带渲染,可以看到车辆运行
|
||||
|
||||
---
|
||||
|
||||
## 🔍 性能监控
|
||||
|
||||
所有版本都内置了性能统计,运行时会显示:
|
||||
```
|
||||
Step 100: FPS = 45.23, 车辆数 = 51, 平均步时间 = 22.10ms
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## ⚙️ 高级优化选项
|
||||
|
||||
### 调整激光雷达缓存频率
|
||||
|
||||
编辑 `run_multiagent_env_fast.py`:
|
||||
```python
|
||||
env.lidar_cache_interval = 3 # 改为5可进一步提速(但观测会更旧)
|
||||
```
|
||||
|
||||
### 调整并行进程数
|
||||
|
||||
编辑 `run_multiagent_env_parallel.py`:
|
||||
```python
|
||||
num_workers = 10 # 改为更少的进程数(如果内存不足)
|
||||
```
|
||||
|
||||
### 进一步减少激光束
|
||||
|
||||
编辑 `scenario_env.py` 的 `_get_all_obs()` 函数:
|
||||
```python
|
||||
lidar = self.engine.get_sensor("lidar").perceive(
|
||||
num_lasers=20, # 从40进一步减少到20
|
||||
distance=20, # 从30减少到20米
|
||||
...
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🎓 为什么CPU利用率低?
|
||||
|
||||
### 原因分析:
|
||||
|
||||
1. **单线程瓶颈**
|
||||
- Python GIL限制
|
||||
- MetaDrive主循环是单线程的
|
||||
- 即使有10个核心,也只用1个
|
||||
|
||||
2. **I/O等待**
|
||||
- 等待渲染完成(如果开启)
|
||||
- 等待磁盘读取数据
|
||||
|
||||
3. **计算不均衡**
|
||||
- 某些计算很重(激光雷达),某些很轻
|
||||
- CPU在重计算之间有空闲
|
||||
|
||||
### 解决方案:
|
||||
|
||||
✅ **已实现:** 多进程并行(`run_multiagent_env_parallel.py`)
|
||||
- 每个进程占用1个核心
|
||||
- 10个进程可充分利用10核CPU
|
||||
- CPU利用率可达90-100%
|
||||
|
||||
---
|
||||
|
||||
## 📚 相关文档
|
||||
|
||||
- `PERFORMANCE_OPTIMIZATION.md` - 详细的性能优化指南
|
||||
- `GPU_ACCELERATION.md` - GPU加速的完整说明
|
||||
|
||||
---
|
||||
|
||||
## ❓ 常见问题
|
||||
|
||||
### Q: 为什么关闭渲染后FPS还是只有20?
|
||||
A: 主要瓶颈是激光雷达计算,不是渲染。请使用 `run_multiagent_env_fast.py`。
|
||||
|
||||
### Q: GPU能加速训练吗?
|
||||
A: 环境模拟在CPU,但策略网络可以在GPU上训练。
|
||||
|
||||
### Q: 如何最大化CPU利用率?
|
||||
A: 使用 `run_multiagent_env_parallel.py` 多进程版本。
|
||||
|
||||
### Q: 会影响观测精度吗?
|
||||
A: 激光束减少会略微降低精度,但实践中影响很小。缓存会让观测滞后1-2帧。
|
||||
|
||||
### Q: 如何恢复原始配置?
|
||||
A: 使用 `run_multiagent_env_visual.py` 或修改配置文件中的参数。
|
||||
|
||||
---
|
||||
|
||||
## 🚦 下一步
|
||||
|
||||
1. 先测试 `run_multiagent_env_fast.py`,验证性能提升
|
||||
2. 如果满意,用于日常训练
|
||||
3. 需要大规模训练时,使用 `run_multiagent_env_parallel.py`
|
||||
4. 考虑将策略网络迁移到GPU
|
||||
|
||||
祝训练顺利!🎉
|
||||
|
||||
BIN
Env/__pycache__/expert_replay_env.cpython-313.pyc
Normal file
BIN
Env/__pycache__/expert_replay_env.cpython-313.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/expert_replay_env.cpython-39.pyc
Normal file
BIN
Env/__pycache__/expert_replay_env.cpython-39.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/expert_replay_policy.cpython-310.pyc
Normal file
BIN
Env/__pycache__/expert_replay_policy.cpython-310.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/inverse_dynamics.cpython-313.pyc
Normal file
BIN
Env/__pycache__/inverse_dynamics.cpython-313.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/inverse_dynamics.cpython-39.pyc
Normal file
BIN
Env/__pycache__/inverse_dynamics.cpython-39.pyc
Normal file
Binary file not shown.
Binary file not shown.
BIN
Env/__pycache__/logger_utils.cpython-39.pyc
Normal file
BIN
Env/__pycache__/logger_utils.cpython-39.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/replay_policy.cpython-39.pyc
Normal file
BIN
Env/__pycache__/replay_policy.cpython-39.pyc
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Env/__pycache__/scenario_env.cpython-39.pyc
Normal file
BIN
Env/__pycache__/scenario_env.cpython-39.pyc
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Env/__pycache__/simple_idm_policy.cpython-39.pyc
Normal file
BIN
Env/__pycache__/simple_idm_policy.cpython-39.pyc
Normal file
Binary file not shown.
194
Env/bc_ego_replay_env.py
Normal file
194
Env/bc_ego_replay_env.py
Normal file
@@ -0,0 +1,194 @@
|
||||
"""
|
||||
Single-agent BC evaluation environment: only ego (SDC) is controlled by the policy;
|
||||
other vehicles are replayed from expert trajectories (same as data collection).
|
||||
"""
|
||||
import numpy as np
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
from Env.hbbc_background_policy import HBBCBackgroundController
|
||||
|
||||
|
||||
class BCEgoReplayEnv(ExpertReplayEnv):
|
||||
"""
|
||||
For single-agent BC evaluation: controlled_agents exposes only SDC (default_agent).
|
||||
Other vehicles are still spawned and replayed by expert; internally we keep them
|
||||
in _replay_agents so step() can update them.
|
||||
"""
|
||||
|
||||
def reset(self, seed=None):
|
||||
obs = super().reset(seed=seed)
|
||||
self.enable_hbbc_background = bool(self.config.get("enable_hbbc_background", False))
|
||||
self.hbbc_controller = None
|
||||
self._hbbc_runtime_logged = False
|
||||
if self.enable_hbbc_background:
|
||||
self.hbbc_controller = HBBCBackgroundController(
|
||||
model_path=self.config.get("hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||
device=self.config.get("hbbc_inference_device", "cpu"),
|
||||
latent_mode=self.config.get("hbbc_latent_mode", "per_vehicle_fixed"),
|
||||
latent_json_path=self.config.get("hbbc_latent_json_path"),
|
||||
seed=int(self.config.get("seed", 0)),
|
||||
dt=float(self.config.get("hbbc_dt", 0.1)),
|
||||
)
|
||||
self.hbbc_controller.reset_episode()
|
||||
# Expose only SDC as the controlled agent for the evaluator
|
||||
self._replay_agents = dict(self.controlled_agents)
|
||||
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||
self.controlled_agents = {self.sdc_agent_id: self.sdc_vehicle}
|
||||
self.controlled_agent_ids = [self.sdc_agent_id]
|
||||
else:
|
||||
self.controlled_agents = {}
|
||||
self.controlled_agent_ids = []
|
||||
return self._get_all_obs()
|
||||
|
||||
def _get_all_obs(self):
|
||||
"""Return only ego (SDC) observation so evaluator has a single agent."""
|
||||
if not self.controlled_agents or self.sdc_vehicle is None:
|
||||
return {}
|
||||
obs = self._obs_for_vehicle(self.sdc_vehicle, exclude_agent_id=self.sdc_agent_id)
|
||||
return {self.sdc_agent_id: obs}
|
||||
|
||||
def step(self, action_dict=None):
|
||||
self.round += 1
|
||||
expert_actions = {}
|
||||
agents_to_remove = []
|
||||
|
||||
# SDC: use policy action if provided, else expert replay
|
||||
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||
policy_action = None
|
||||
if action_dict and self.sdc_agent_id in action_dict:
|
||||
policy_action = np.asarray(action_dict[self.sdc_agent_id], dtype=np.float64)
|
||||
next_step = self.round
|
||||
curr_step = self.round - 1
|
||||
if next_step < len(self.sdc_track["state"]["position"]) and self.sdc_track["state"]["valid"][next_step]:
|
||||
curr_state = {
|
||||
"position": self.sdc_track["state"]["position"][curr_step],
|
||||
"heading": self.sdc_track["state"]["heading"][curr_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][curr_step],
|
||||
}
|
||||
if policy_action is not None:
|
||||
next_state = self.inverse_dynamics.apply_action(curr_state, policy_action, dt=0.1)
|
||||
expert_actions[self.sdc_agent_id] = policy_action
|
||||
else:
|
||||
next_state = {
|
||||
"position": self.sdc_track["state"]["position"][next_step],
|
||||
"heading": self.sdc_track["state"]["heading"][next_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][next_step],
|
||||
}
|
||||
action, _ = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1)
|
||||
expert_actions[self.sdc_agent_id] = action
|
||||
self.sdc_vehicle.set_position(next_state["position"])
|
||||
self.sdc_vehicle.set_heading_theta(next_state["heading"])
|
||||
self.sdc_vehicle.set_velocity(next_state["velocity"])
|
||||
self.sdc_vehicle.last_expert_action = expert_actions[self.sdc_agent_id]
|
||||
|
||||
# Replay other vehicles: restore full controlled_agents for internal logic
|
||||
self.controlled_agents = dict(self._replay_agents)
|
||||
self.controlled_agent_ids = list(self.controlled_agents.keys())
|
||||
hbbc_batch = []
|
||||
hbbc_curr_states = {}
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
track = vehicle.expert_track
|
||||
next_step = self.round
|
||||
if next_step >= len(track["state"]["position"]):
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
if not track["state"]["valid"][next_step]:
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
if self.enable_hbbc_background and self.hbbc_controller is not None:
|
||||
# HBBC autonomous rollout: use vehicle's own previous-step state
|
||||
curr_state = {
|
||||
"position": np.asarray(vehicle.position, dtype=np.float64),
|
||||
"heading": float(vehicle.heading_theta),
|
||||
"velocity": np.asarray(vehicle.velocity, dtype=np.float64),
|
||||
}
|
||||
object_id = str(getattr(vehicle, "original_id", agent_id))
|
||||
hbbc_batch.append((agent_id, vehicle, object_id, agent_id))
|
||||
hbbc_curr_states[agent_id] = curr_state
|
||||
else:
|
||||
curr_step = self.round - 1
|
||||
curr_state = {
|
||||
"position": track["state"]["position"][curr_step],
|
||||
"heading": track["state"]["heading"][curr_step],
|
||||
"velocity": track["state"]["velocity"][curr_step],
|
||||
}
|
||||
next_state = {
|
||||
"position": track["state"]["position"][next_step],
|
||||
"heading": track["state"]["heading"][next_step],
|
||||
"velocity": track["state"]["velocity"][next_step],
|
||||
}
|
||||
action, _ = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1)
|
||||
expert_actions[agent_id] = action
|
||||
vehicle.set_position(next_state["position"])
|
||||
vehicle.set_heading_theta(next_state["heading"])
|
||||
vehicle.set_velocity(next_state["velocity"])
|
||||
vehicle.last_expert_action = action
|
||||
|
||||
if hbbc_batch and self.hbbc_controller is not None:
|
||||
hbbc_actions = self.hbbc_controller.infer_actions(hbbc_batch)
|
||||
if not self._hbbc_runtime_logged:
|
||||
print(f"[HBBC] background policy active, current dynamic agents: {len(hbbc_batch)}")
|
||||
self._hbbc_runtime_logged = True
|
||||
for agent_id, _, _, _ in hbbc_batch:
|
||||
curr_state = hbbc_curr_states[agent_id]
|
||||
action = hbbc_actions[agent_id]
|
||||
next_state = self.inverse_dynamics.apply_action(curr_state, action, dt=0.1)
|
||||
expert_actions[agent_id] = action
|
||||
vehicle = self.controlled_agents[agent_id]
|
||||
vehicle.set_position(next_state["position"])
|
||||
vehicle.set_heading_theta(next_state["heading"])
|
||||
vehicle.set_velocity(next_state["velocity"])
|
||||
try:
|
||||
vehicle.last_current_action.append(action)
|
||||
except Exception:
|
||||
pass
|
||||
vehicle.last_expert_action = action
|
||||
for agent_id in agents_to_remove:
|
||||
vehicle = self.controlled_agents[agent_id]
|
||||
self.controlled_agents.pop(agent_id)
|
||||
self.controlled_agent_ids.remove(agent_id)
|
||||
self.engine.agent_manager.active_agents.pop(agent_id, None)
|
||||
self.engine.clear_objects([vehicle.id])
|
||||
if self.hbbc_controller is not None:
|
||||
self.hbbc_controller.remove_vehicle(agent_id)
|
||||
self.engine.taskMgr.step()
|
||||
self._spawn_controlled_agents()
|
||||
self._update_background_vehicles()
|
||||
self._replay_agents = dict(self.controlled_agents)
|
||||
# Expose only SDC again
|
||||
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||
self.controlled_agents = {self.sdc_agent_id: self.sdc_vehicle}
|
||||
self.controlled_agent_ids = [self.sdc_agent_id]
|
||||
else:
|
||||
self.controlled_agents = {}
|
||||
self.controlled_agent_ids = []
|
||||
|
||||
obs = self._get_all_obs()
|
||||
rewards = {}
|
||||
infos = {aid: {"expert_action": expert_actions.get(aid, np.zeros(2))} for aid in self.controlled_agents}
|
||||
if self.sdc_agent_id in self.controlled_agents and self.sdc_vehicle is not None:
|
||||
speed_coef = float(self.config.get("reward_speed_coef", 0.05))
|
||||
collision_distance = float(self.config.get("collision_distance", 6.0))
|
||||
collision_penalty = float(self.config.get("collision_penalty", 100.0))
|
||||
speed = float(np.linalg.norm(self.sdc_vehicle.velocity))
|
||||
r_speed = speed_coef * speed
|
||||
min_dist = float("inf")
|
||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
if other_id == self.sdc_agent_id:
|
||||
continue
|
||||
try:
|
||||
d = float(np.linalg.norm(self.sdc_vehicle.position - other_vehicle.position))
|
||||
min_dist = min(min_dist, d)
|
||||
except Exception:
|
||||
continue
|
||||
near_collision = min_dist < collision_distance
|
||||
r_collision = -collision_penalty if near_collision else 0.0
|
||||
rewards[self.sdc_agent_id] = r_speed + r_collision
|
||||
infos[self.sdc_agent_id].update(
|
||||
near_collision=near_collision,
|
||||
min_dist=min_dist if np.isfinite(min_dist) else None,
|
||||
r_speed=r_speed,
|
||||
r_collision=r_collision,
|
||||
)
|
||||
dones = {aid: False for aid in self.controlled_agents}
|
||||
dones["__all__"] = self.round >= self.config["horizon"] or (len(self._replay_agents) == 0 and self.round > 190)
|
||||
return obs, rewards, dones, infos
|
||||
246
Env/bc_env.py
Normal file
246
Env/bc_env.py
Normal file
@@ -0,0 +1,246 @@
|
||||
from Env.scenario_env import MultiAgentScenarioEnv
|
||||
from Env.hbbc_background_policy import HBBCBackgroundController
|
||||
from Env.utils import filter_traffic_tracks_to_birth_lists
|
||||
from metadrive.component.vehicle.vehicle_type import DefaultVehicle
|
||||
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)
|
||||
|
||||
Spawns background (static) vehicles so that observation distribution matches expert data collection:
|
||||
expert data is generated with ExpertReplayEnv which includes bg_* in active_agents, so the policy
|
||||
was trained on obs that can include those neighbors. Demo should use the same scene for consistency.
|
||||
"""
|
||||
def _init_hbbc_background(self):
|
||||
self.enable_hbbc_background = bool(self.config.get("enable_hbbc_background", False))
|
||||
self.hbbc_dynamic_agents = {}
|
||||
self._spawned_dynamic_bg_ids = set()
|
||||
self.hbbc_controller = None
|
||||
if not self.enable_hbbc_background:
|
||||
return
|
||||
self.hbbc_controller = HBBCBackgroundController(
|
||||
model_path=self.config.get("hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||
device=self.config.get("hbbc_inference_device", "cpu"),
|
||||
latent_mode=self.config.get("hbbc_latent_mode", "per_vehicle_fixed"),
|
||||
latent_json_path=self.config.get("hbbc_latent_json_path"),
|
||||
seed=int(self.config.get("seed", 0)),
|
||||
dt=float(self.config.get("hbbc_dt", 0.1)),
|
||||
)
|
||||
self.hbbc_controller.reset_episode()
|
||||
|
||||
def _move_excess_controlled_to_hbbc_background(self):
|
||||
if not self.enable_hbbc_background:
|
||||
return
|
||||
keep_n = int(self.config.get("num_controlled_agents", 0))
|
||||
keep_n = max(0, keep_n)
|
||||
ordered_ids = list(self.controlled_agents.keys())
|
||||
keep_ids = set(ordered_ids[:keep_n])
|
||||
move_ids = [aid for aid in ordered_ids if aid not in keep_ids]
|
||||
for aid in move_ids:
|
||||
self.hbbc_dynamic_agents[aid] = self.controlled_agents[aid]
|
||||
self.controlled_agents.pop(aid, None)
|
||||
if aid in self.controlled_agent_ids:
|
||||
self.controlled_agent_ids.remove(aid)
|
||||
self._spawned_dynamic_bg_ids.update(move_ids)
|
||||
|
||||
def _apply_hbbc_before_step(self):
|
||||
if not self.enable_hbbc_background or not self.hbbc_dynamic_agents:
|
||||
return
|
||||
batch = []
|
||||
for aid, vehicle in self.hbbc_dynamic_agents.items():
|
||||
object_id = getattr(vehicle, "original_id", None) or aid.replace("controlled_", "", 1)
|
||||
batch.append((aid, vehicle, str(object_id) if object_id is not None else None, aid))
|
||||
actions = self.hbbc_controller.infer_actions(batch)
|
||||
for aid, vehicle in self.hbbc_dynamic_agents.items():
|
||||
action = actions.get(aid, np.zeros(2, dtype=np.float32))
|
||||
vehicle.before_step(action)
|
||||
|
||||
def _apply_hbbc_after_step(self):
|
||||
if not self.enable_hbbc_background:
|
||||
return
|
||||
for vehicle in self.hbbc_dynamic_agents.values():
|
||||
vehicle.after_step()
|
||||
|
||||
def reset(self, seed=None):
|
||||
self._init_hbbc_background()
|
||||
# Clear background vehicles from previous episode so engine.reset() passes _object_clean_check
|
||||
if getattr(self, "engine", None) is not None:
|
||||
ids_bg = [
|
||||
oid for oid, obj in self.engine.get_objects().items()
|
||||
if (getattr(obj, "name", None) or getattr(obj, "id", None) or "").startswith(("bg_", "controlled_"))
|
||||
]
|
||||
if ids_bg:
|
||||
self.engine.clear_objects(ids_bg, force_destroy=True)
|
||||
for aid in list(self.engine.agent_manager.active_agents.keys()):
|
||||
if aid.startswith("bg_") or aid.startswith("controlled_"):
|
||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||
obs = super().reset(seed=seed)
|
||||
self._move_excess_controlled_to_hbbc_background()
|
||||
self._spawn_background_vehicles()
|
||||
return self._get_all_obs()
|
||||
|
||||
def _build_birth_lists_from_traffic(self):
|
||||
"""Same lane/static filter as expert data; return background_vehicles so we spawn them (match training obs)."""
|
||||
car_birth_info_list, background_vehicles, obj_to_clean, stats = filter_traffic_tracks_to_birth_lists(
|
||||
self.engine.traffic_manager.current_traffic_data,
|
||||
self.engine.traffic_manager.sdc_scenario_id,
|
||||
self.engine.map_manager,
|
||||
return_stats=True,
|
||||
)
|
||||
if stats["n_controlled"] == 0 and stats["n_total"] > 0:
|
||||
print(
|
||||
"[BCScenarioEnv] 0 controlled agents: total_vehicles={}, off_lane={}, static={}, no_valid={}.".format(
|
||||
stats["n_total"],
|
||||
stats["n_off_lane"],
|
||||
stats["n_static"],
|
||||
stats["n_no_valid"],
|
||||
)
|
||||
)
|
||||
return car_birth_info_list, background_vehicles, obj_to_clean
|
||||
|
||||
def _spawn_background_vehicles(self):
|
||||
"""Spawn all static background vehicles once at reset (no show_time filter; same as ExpertReplayEnv)."""
|
||||
for sid, car in self.background_vehicles.items():
|
||||
bg_id = f"bg_{car['id']}"
|
||||
if bg_id in self.engine.agent_manager.active_agents:
|
||||
continue
|
||||
vehicle_config = {}
|
||||
if "length" in car and "width" in car:
|
||||
vehicle_config = {"length": car["length"], "width": car["width"]}
|
||||
v = self.engine.spawn_object(
|
||||
DefaultVehicle,
|
||||
name=bg_id,
|
||||
vehicle_config=vehicle_config,
|
||||
position=car["begin"],
|
||||
heading=car["heading"],
|
||||
)
|
||||
v.set_velocity([0, 0])
|
||||
self.engine.agent_manager.active_agents[bg_id] = v
|
||||
v.valid_mask = car.get("valid")
|
||||
v.start_t = car.get("show_time")
|
||||
|
||||
def _update_background_vehicles(self):
|
||||
# Static vehicles are spawned once at init and never removed.
|
||||
pass
|
||||
|
||||
def step(self, action_dict):
|
||||
if action_dict is None:
|
||||
action_dict = {}
|
||||
self.round += 1
|
||||
for agent_id, action in action_dict.items():
|
||||
if agent_id in self.controlled_agents:
|
||||
self.controlled_agents[agent_id].before_step(action)
|
||||
self._apply_hbbc_before_step()
|
||||
self.engine.step()
|
||||
self.engine.after_step()
|
||||
for agent_id in action_dict:
|
||||
if agent_id in self.controlled_agents:
|
||||
self.controlled_agents[agent_id].after_step()
|
||||
self._spawn_controlled_agents()
|
||||
self._move_excess_controlled_to_hbbc_background()
|
||||
self._apply_hbbc_after_step()
|
||||
self._update_background_vehicles()
|
||||
obs = self._get_all_obs()
|
||||
|
||||
# Reward shaping for evaluation/rollout monitoring (BC training itself doesn't use env reward).
|
||||
speed_coef = float(self.config.get("reward_speed_coef", 0.05))
|
||||
collision_distance = float(self.config.get("collision_distance", 6.0))
|
||||
collision_penalty = float(self.config.get("collision_penalty", 100.0))
|
||||
|
||||
# Pre-collect all active vehicles (includes background vehicles).
|
||||
active_agents = list(self.engine.agent_manager.active_agents.items())
|
||||
|
||||
rewards = {}
|
||||
infos = {}
|
||||
for aid, vehicle in self.controlled_agents.items():
|
||||
# Speed reward
|
||||
speed = getattr(vehicle, "speed", None)
|
||||
if speed is None:
|
||||
speed = float(np.linalg.norm(vehicle.velocity))
|
||||
r_speed = speed_coef * float(speed)
|
||||
|
||||
# Near-collision penalty (distance-based, simulator-agnostic)
|
||||
min_dist = float("inf")
|
||||
for other_id, other_vehicle in active_agents:
|
||||
if other_id == aid:
|
||||
continue
|
||||
try:
|
||||
dist = float(np.linalg.norm(vehicle.position - other_vehicle.position))
|
||||
except Exception:
|
||||
continue
|
||||
if dist < min_dist:
|
||||
min_dist = dist
|
||||
|
||||
near_collision = bool(min_dist < collision_distance)
|
||||
r_collision = -collision_penalty if near_collision else 0.0
|
||||
|
||||
rewards[aid] = float(r_speed + r_collision)
|
||||
infos[aid] = {
|
||||
"near_collision": near_collision,
|
||||
"min_dist": (min_dist if np.isfinite(min_dist) else None),
|
||||
"r_speed": float(r_speed),
|
||||
"r_collision": float(r_collision),
|
||||
}
|
||||
dones = {aid: False for aid in self.controlled_agents}
|
||||
dones["__all__"] = self.episode_step >= self.config["horizon"]
|
||||
return obs, rewards, dones, infos
|
||||
|
||||
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
|
||||
@@ -1,116 +0,0 @@
|
||||
"""
|
||||
日志记录功能示例
|
||||
演示如何在自定义脚本中使用日志功能
|
||||
"""
|
||||
from logger_utils import setup_logger
|
||||
from datetime import datetime
|
||||
import time
|
||||
|
||||
def example_without_logging():
|
||||
"""示例1:不使用日志"""
|
||||
print("=" * 60)
|
||||
print("示例1:普通输出(不记录日志)")
|
||||
print("=" * 60)
|
||||
|
||||
print("这是普通的print输出")
|
||||
print("只会显示在终端")
|
||||
print("不会保存到文件")
|
||||
print()
|
||||
|
||||
|
||||
def example_with_logging():
|
||||
"""示例2:使用日志记录"""
|
||||
print("=" * 60)
|
||||
print("示例2:使用日志记录")
|
||||
print("=" * 60)
|
||||
|
||||
# 使用with语句,自动管理日志文件
|
||||
with setup_logger(log_file="example_demo.log", log_dir="logs"):
|
||||
print("✅ 这条消息会同时输出到终端和文件")
|
||||
print("✅ 运行一些计算...")
|
||||
|
||||
for i in range(5):
|
||||
print(f" 步骤 {i+1}/5: 处理中...")
|
||||
time.sleep(0.1)
|
||||
|
||||
print("✅ 计算完成!")
|
||||
|
||||
print("日志文件已关闭")
|
||||
print()
|
||||
|
||||
|
||||
def example_custom_filename():
|
||||
"""示例3:使用时间戳命名"""
|
||||
print("=" * 60)
|
||||
print("示例3:自动生成时间戳文件名")
|
||||
print("=" * 60)
|
||||
|
||||
# log_file=None 会自动生成时间戳文件名
|
||||
with setup_logger(log_file=None, log_dir="logs"):
|
||||
print("文件名会自动包含时间戳")
|
||||
print("适合批量实验,避免覆盖")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
def example_append_mode():
|
||||
"""示例4:追加模式"""
|
||||
print("=" * 60)
|
||||
print("示例4:追加到现有文件")
|
||||
print("=" * 60)
|
||||
|
||||
# 第一次写入
|
||||
with setup_logger(log_file="append_test.log", log_dir="logs", mode='w'):
|
||||
print("第一次写入:这会覆盖文件")
|
||||
|
||||
# 第二次写入(追加)
|
||||
with setup_logger(log_file="append_test.log", log_dir="logs", mode='a'):
|
||||
print("第二次写入:这会追加到文件末尾")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
def example_complex_output():
|
||||
"""示例5:复杂输出(包含颜色、格式)"""
|
||||
print("=" * 60)
|
||||
print("示例5:复杂输出格式")
|
||||
print("=" * 60)
|
||||
|
||||
with setup_logger(log_file="complex_output.log", log_dir="logs"):
|
||||
# 模拟多种输出格式
|
||||
print("\n📊 实验统计:")
|
||||
print(" - 实验名称:车道过滤测试")
|
||||
print(" - 开始时间:", datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
|
||||
print(" - 车辆总数:51")
|
||||
print(" - 过滤后:45")
|
||||
print("\n🚦 红绿灯检测:")
|
||||
print(" ✅ 方法1成功:3辆")
|
||||
print(" ✅ 方法2成功:2辆")
|
||||
print(" ⚠️ 未检测到:40辆")
|
||||
print("\n" + "="*50)
|
||||
print("实验完成!")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
def main():
|
||||
"""运行所有示例"""
|
||||
print("\n" + "🎯 " + "="*56)
|
||||
print("日志记录功能完整示例")
|
||||
print("="*60 + "\n")
|
||||
|
||||
example_without_logging()
|
||||
example_with_logging()
|
||||
example_custom_filename()
|
||||
example_append_mode()
|
||||
example_complex_output()
|
||||
|
||||
print("="*60)
|
||||
print("✅ 所有示例运行完成!")
|
||||
print("📁 查看日志文件:ls -lh logs/")
|
||||
print("="*60)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
345
Env/expert_replay_env.py
Normal file
345
Env/expert_replay_env.py
Normal file
@@ -0,0 +1,345 @@
|
||||
import logging
|
||||
import numpy as np
|
||||
from collections import defaultdict
|
||||
from metadrive.component.vehicle.vehicle_type import DefaultVehicle
|
||||
from metadrive.type import MetaDriveType
|
||||
from Env.scenario_env import MultiAgentScenarioEnv, PolicyVehicle
|
||||
from Env.inverse_dynamics import InverseDynamics
|
||||
|
||||
class ExpertReplayEnv(MultiAgentScenarioEnv):
|
||||
def __init__(self, config=None):
|
||||
# Allow passing config without agent2policy since we don't use policies for replay
|
||||
if config is None:
|
||||
config = {}
|
||||
# Ensure we don't simulate physics for the controlled agents in the traditional sense
|
||||
# but we still need the engine to run
|
||||
super().__init__(config, agent2policy={})
|
||||
self.inverse_dynamics = InverseDynamics()
|
||||
self.expert_tracks = {}
|
||||
# Replay SDC/ego ("default_agent" in MetaDrive) as well; otherwise it will keep default action=0 and look stuck.
|
||||
self.replay_sdc = self.config.get("replay_sdc", True)
|
||||
self.sdc_track = None
|
||||
self.sdc_vehicle = None
|
||||
self.sdc_agent_id = "default_agent"
|
||||
|
||||
def reset(self, seed=None):
|
||||
self.round = 0
|
||||
if self.logger is None:
|
||||
from metadrive.engine.logger import get_logger, set_log_level
|
||||
self.logger = get_logger()
|
||||
log_level = self.config.get("log_level", logging.INFO)
|
||||
set_log_level(log_level)
|
||||
|
||||
self.lazy_init()
|
||||
self._reset_global_seed(seed)
|
||||
if self.engine is None:
|
||||
raise ValueError("Broken MetaDrive instance.")
|
||||
|
||||
self.background_vehicles = {}
|
||||
self.expert_tracks = {}
|
||||
self.sdc_track = None
|
||||
self.sdc_vehicle = None
|
||||
|
||||
# 在加载新场景前,必须清除上一轮通过 spawn_object 生成的物体,否则 engine.reset() 内 _object_clean_check 会报错
|
||||
# 从 engine 当前对象中按名称筛选(与 manager 无关的对象需在此清理),并强制销毁
|
||||
ids_to_clear = []
|
||||
for oid, obj in self.engine.get_objects().items():
|
||||
name = getattr(obj, "name", None) or getattr(obj, "id", None)
|
||||
if name and (str(name).startswith("controlled_") or str(name).startswith("bg_")):
|
||||
ids_to_clear.append(oid)
|
||||
if ids_to_clear:
|
||||
self.engine.clear_objects(ids_to_clear, force_destroy=True)
|
||||
self.controlled_agents.clear()
|
||||
self.controlled_agent_ids.clear()
|
||||
for aid in list(self.engine.agent_manager.active_agents.keys()):
|
||||
if aid.startswith("bg_") or aid.startswith("controlled_"):
|
||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||
|
||||
if self.replay_sdc and hasattr(self.engine, "traffic_manager"):
|
||||
sdc_sid = self.engine.traffic_manager.sdc_scenario_id
|
||||
self.sdc_track = self.engine.traffic_manager.current_traffic_data.get(sdc_sid, None)
|
||||
|
||||
from Env.utils import filter_traffic_tracks_to_birth_lists
|
||||
traffic_data = self.engine.traffic_manager.current_traffic_data
|
||||
car_birth_info_list, self.background_vehicles, obj_to_clean = filter_traffic_tracks_to_birth_lists(
|
||||
traffic_data,
|
||||
self.engine.traffic_manager.sdc_scenario_id,
|
||||
self.engine.map_manager,
|
||||
)
|
||||
for entry in car_birth_info_list:
|
||||
sid = entry["scenario_id"]
|
||||
if sid in traffic_data:
|
||||
self.expert_tracks[sid] = traffic_data[sid]
|
||||
self.car_birth_info_list = car_birth_info_list
|
||||
for scenario_id in obj_to_clean:
|
||||
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||
|
||||
self.engine.reset()
|
||||
self.reset_sensors()
|
||||
self.engine.taskMgr.step()
|
||||
|
||||
self.lanes = self.engine.map_manager.current_map.road_network.graph
|
||||
|
||||
if self.top_down_renderer is not None:
|
||||
self.top_down_renderer.clear()
|
||||
self.engine.top_down_renderer = None
|
||||
|
||||
self.dones = {}
|
||||
self.episode_rewards = defaultdict(float)
|
||||
self.episode_lengths = defaultdict(int)
|
||||
|
||||
self.controlled_agents.clear()
|
||||
self.controlled_agent_ids.clear()
|
||||
|
||||
# We skip calling super().reset() to avoid double reset
|
||||
# But we need to ensure ScenarioEnv-specific setup is done if any.
|
||||
# ScenarioEnv.reset() basically does engine.reset() and some cleanup.
|
||||
# We covered most of it.
|
||||
|
||||
self._spawn_controlled_agents()
|
||||
self._spawn_all_background_vehicles_at_init()
|
||||
|
||||
# Ensure SDC/ego is moved to the correct initial expert state.
|
||||
if self.replay_sdc:
|
||||
self.sdc_vehicle = self.engine.agent_manager.active_agents.get(self.sdc_agent_id, None)
|
||||
if self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||
valid = self.sdc_track["state"]["valid"]
|
||||
t0 = int(np.argmax(valid)) if valid.any() else 0
|
||||
pos0 = self.sdc_track["state"]["position"][t0]
|
||||
heading0 = self.sdc_track["state"]["heading"][t0]
|
||||
vel0 = self.sdc_track["state"]["velocity"][t0]
|
||||
self.sdc_vehicle.set_position(pos0)
|
||||
self.sdc_vehicle.set_heading_theta(heading0)
|
||||
self.sdc_vehicle.set_velocity(vel0)
|
||||
|
||||
return self._get_all_obs()
|
||||
|
||||
def _spawn_all_background_vehicles_at_init(self):
|
||||
"""Spawn all static background vehicles once at reset (no show_time filter; no removal by valid)."""
|
||||
for sid, car in self.background_vehicles.items():
|
||||
bg_id = f"bg_{car['id']}"
|
||||
if bg_id in self.engine.agent_manager.active_agents:
|
||||
continue
|
||||
vehicle_config = {}
|
||||
if 'length' in car and 'width' in car:
|
||||
vehicle_config = {
|
||||
"length": car['length'],
|
||||
"width": car['width']
|
||||
}
|
||||
v = self.engine.spawn_object(
|
||||
DefaultVehicle,
|
||||
name=bg_id,
|
||||
vehicle_config=vehicle_config,
|
||||
position=car['begin'],
|
||||
heading=car['heading']
|
||||
)
|
||||
v.set_velocity([0, 0])
|
||||
self.engine.agent_manager.active_agents[bg_id] = v
|
||||
v.valid_mask = car.get('valid')
|
||||
v.start_t = car['show_time']
|
||||
|
||||
def _update_background_vehicles(self):
|
||||
# Static vehicles are spawned once at init and never removed (no spawn/remove by show_time or valid).
|
||||
pass
|
||||
|
||||
def _spawn_controlled_agents(self):
|
||||
for car in self.car_birth_info_list:
|
||||
if car['show_time'] == self.round:
|
||||
agent_id = f"controlled_{car['id']}"
|
||||
|
||||
# Check if we already have this agent (shouldn't happen with unique IDs but safety check)
|
||||
if agent_id in self.controlled_agents:
|
||||
continue
|
||||
|
||||
# Handling ID flickering / merging
|
||||
# If this ID is new, check if there's an existing agent very close to its start position
|
||||
# that just disappeared? (Not implemented here, complex logic)
|
||||
# But we can check if there's an overlap with existing agents?
|
||||
# For now, just spawn.
|
||||
|
||||
# Read vehicle type/size if available
|
||||
vehicle_config = {}
|
||||
if 'length' in car and 'width' in car:
|
||||
vehicle_config = {
|
||||
"length": car['length'],
|
||||
"width": car['width']
|
||||
}
|
||||
|
||||
vehicle = self.engine.spawn_object(
|
||||
PolicyVehicle,
|
||||
name=agent_id,
|
||||
vehicle_config=vehicle_config,
|
||||
position=car['begin'],
|
||||
heading=car['heading']
|
||||
)
|
||||
vehicle.reset(position=car['begin'], heading=car['heading'])
|
||||
|
||||
# We don't set policy or destination in the same way, or maybe we do for compatibility
|
||||
vehicle.set_destination(car['end'])
|
||||
|
||||
# Store extra info for replay
|
||||
vehicle.expert_track = self.expert_tracks[car['scenario_id']]
|
||||
vehicle.original_id = car['id']
|
||||
|
||||
self.controlled_agents[agent_id] = vehicle
|
||||
self.controlled_agent_ids.append(agent_id)
|
||||
self.engine.agent_manager.active_agents[agent_id] = vehicle
|
||||
|
||||
def step(self, action_dict=None):
|
||||
# We ignore input action_dict for the purpose of controlling agents
|
||||
# Instead, we calculate what the action *should* be
|
||||
|
||||
self.round += 1
|
||||
expert_actions = {}
|
||||
|
||||
# 1. Update state of all controlled agents to the current timestep (self.round)
|
||||
# and compute action from (self.round-1) to (self.round).
|
||||
# Wait, usually step() moves T -> T+1.
|
||||
# Current state is T. We want to move to T+1.
|
||||
# So we need state at T and T+1.
|
||||
|
||||
# Identify agents that are done (valid=0 at T+1 or T+1 >= length)
|
||||
agents_to_remove = []
|
||||
|
||||
# Update SDC/ego first (otherwise it will stay still with default action=0)
|
||||
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||
next_step = self.round
|
||||
curr_step = self.round - 1
|
||||
if next_step < len(self.sdc_track["state"]["position"]) and self.sdc_track["state"]["valid"][next_step]:
|
||||
curr_state = {
|
||||
"position": self.sdc_track["state"]["position"][curr_step],
|
||||
"heading": self.sdc_track["state"]["heading"][curr_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][curr_step],
|
||||
}
|
||||
next_state = {
|
||||
"position": self.sdc_track["state"]["position"][next_step],
|
||||
"heading": self.sdc_track["state"]["heading"][next_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][next_step],
|
||||
}
|
||||
action, _ = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1)
|
||||
expert_actions[self.sdc_agent_id] = action
|
||||
self.sdc_vehicle.set_position(next_state["position"])
|
||||
self.sdc_vehicle.set_heading_theta(next_state["heading"])
|
||||
self.sdc_vehicle.set_velocity(next_state["velocity"])
|
||||
self.sdc_vehicle.last_expert_action = action
|
||||
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
track = vehicle.expert_track
|
||||
# current_step = self.round - 1 # Since we incremented at start
|
||||
# But vehicle is currently at state corresponding to self.round - 1.
|
||||
# We want to move it to self.round.
|
||||
|
||||
# Check bounds
|
||||
next_step = self.round
|
||||
curr_step = self.round - 1
|
||||
|
||||
if next_step >= len(track['state']['position']):
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
|
||||
valid = track['state']['valid'][next_step]
|
||||
if not valid:
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
|
||||
# Get states
|
||||
curr_pos = track['state']['position'][curr_step]
|
||||
next_pos = track['state']['position'][next_step]
|
||||
curr_heading = track['state']['heading'][curr_step]
|
||||
next_heading = track['state']['heading'][next_step]
|
||||
curr_vel = track['state']['velocity'][curr_step]
|
||||
next_vel = track['state']['velocity'][next_step]
|
||||
|
||||
# Prepare state dicts for Inverse Dynamics
|
||||
curr_state = {
|
||||
'position': curr_pos,
|
||||
'heading': curr_heading,
|
||||
'velocity': curr_vel
|
||||
}
|
||||
next_state = {
|
||||
'position': next_pos,
|
||||
'heading': next_heading,
|
||||
'velocity': next_vel
|
||||
}
|
||||
|
||||
# Calculate action
|
||||
action, raw_info = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1) # Waymo is 10Hz?
|
||||
expert_actions[agent_id] = action
|
||||
|
||||
# Force update vehicle state
|
||||
vehicle.set_position(next_pos)
|
||||
vehicle.set_heading_theta(next_heading)
|
||||
vehicle.set_velocity(next_vel)
|
||||
|
||||
# Also record this action in the vehicle for later retrieval if needed
|
||||
vehicle.last_expert_action = action
|
||||
|
||||
# Remove finished agents
|
||||
for agent_id in agents_to_remove:
|
||||
vehicle = self.controlled_agents[agent_id]
|
||||
self.controlled_agents.pop(agent_id)
|
||||
self.controlled_agent_ids.remove(agent_id)
|
||||
self.engine.agent_manager.active_agents.pop(agent_id, None)
|
||||
|
||||
self.engine.clear_objects([vehicle.id])
|
||||
|
||||
# Step physics world to update sensors/collision detection
|
||||
# We don't need full integration, but we need to update the physics world state
|
||||
self.engine.taskMgr.step()
|
||||
|
||||
# Spawn new agents for this turn
|
||||
self._spawn_controlled_agents()
|
||||
self._update_background_vehicles()
|
||||
|
||||
# Get observations
|
||||
obs = self._get_all_obs()
|
||||
|
||||
# Build rewards/dones/infos: include controlled_agents and optionally SDC for data collection
|
||||
all_agent_ids = list(self.controlled_agents.keys())
|
||||
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_agent_id not in all_agent_ids:
|
||||
all_agent_ids = all_agent_ids + [self.sdc_agent_id]
|
||||
rewards = {aid: 0.0 for aid in all_agent_ids}
|
||||
dones = {aid: False for aid in all_agent_ids}
|
||||
dones["__all__"] = (self.round >= self.config["horizon"]) or (len(self.controlled_agents) == 0 and self.round > 190) # Waymo scenarios are usually ~198 steps (20s @ 10Hz) or 90 steps (9s)
|
||||
infos = {aid: {"expert_action": expert_actions.get(aid, np.zeros(2))} for aid in all_agent_ids}
|
||||
|
||||
return obs, rewards, dones, infos
|
||||
|
||||
def _obs_for_vehicle(self, vehicle, exclude_agent_id=None):
|
||||
"""Compute 45-dim obs (ego 5 + 10 neighbors x 4) for a vehicle. exclude_agent_id: do not count as neighbor."""
|
||||
ego_state = [
|
||||
vehicle.position[0], vehicle.position[1],
|
||||
vehicle.velocity[0], vehicle.velocity[1],
|
||||
vehicle.heading_theta
|
||||
]
|
||||
candidates = []
|
||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
if other_id == exclude_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))
|
||||
return np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||
|
||||
def _get_all_obs(self):
|
||||
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||
obs_dict = {}
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
obs_dict[agent_id] = self._obs_for_vehicle(vehicle, exclude_agent_id=agent_id)
|
||||
# Include SDC/ego obs for expert data collection (e.g. single-agent)
|
||||
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||
obs_dict[self.sdc_agent_id] = self._obs_for_vehicle(self.sdc_vehicle, exclude_agent_id=self.sdc_agent_id)
|
||||
return obs_dict
|
||||
69
Env/hbbc_actor_critic.py
Normal file
69
Env/hbbc_actor_critic.py
Normal file
@@ -0,0 +1,69 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def _get_activation(name: str):
|
||||
name = (name or "elu").lower()
|
||||
mapping = {
|
||||
"elu": nn.ELU,
|
||||
"relu": nn.ReLU,
|
||||
"tanh": nn.Tanh,
|
||||
"leakyrelu": nn.LeakyReLU,
|
||||
}
|
||||
if name not in mapping:
|
||||
raise ValueError(f"Unsupported activation: {name}")
|
||||
return mapping[name]()
|
||||
|
||||
|
||||
class ActorCritic(nn.Module):
|
||||
"""Minimal HBBC ActorCritic for inference-only deployment."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_actor_obs=18,
|
||||
num_critic_obs=18,
|
||||
num_actions=2,
|
||||
latent_c_dim=4,
|
||||
latent_eps_dim=6,
|
||||
use_style_latent=True,
|
||||
actor_hidden_dims=None,
|
||||
activation="elu",
|
||||
):
|
||||
super().__init__()
|
||||
_ = num_critic_obs # kept for checkpoint compatibility
|
||||
if actor_hidden_dims is None:
|
||||
actor_hidden_dims = [512, 256, 128]
|
||||
|
||||
act_fn = _get_activation(activation)
|
||||
self.latent_c_dim = int(latent_c_dim)
|
||||
self.latent_eps_dim = int(latent_eps_dim)
|
||||
self.use_style_latent = bool(use_style_latent)
|
||||
|
||||
layers = [nn.Linear(num_actor_obs, actor_hidden_dims[0]), act_fn]
|
||||
for i in range(len(actor_hidden_dims) - 1):
|
||||
layers.append(nn.Linear(actor_hidden_dims[i], actor_hidden_dims[i + 1]))
|
||||
layers.append(_get_activation(activation))
|
||||
self.actor_trunk = nn.Sequential(*layers)
|
||||
self.actor_head = nn.Linear(actor_hidden_dims[-1], num_actions)
|
||||
|
||||
if self.use_style_latent:
|
||||
self.style_trunk = nn.Sequential(
|
||||
nn.Linear(self.latent_eps_dim, 512),
|
||||
_get_activation(activation),
|
||||
nn.Linear(512, 256),
|
||||
_get_activation(activation),
|
||||
nn.Linear(256, 128),
|
||||
_get_activation(activation),
|
||||
)
|
||||
self.style_head = nn.Linear(128, self.latent_eps_dim)
|
||||
self.style_activation = torch.tanh
|
||||
|
||||
def act_inference(self, observations: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_style_latent:
|
||||
obs = observations[..., :-(self.latent_c_dim + self.latent_eps_dim)]
|
||||
eps = observations[..., -self.latent_c_dim - self.latent_eps_dim:-self.latent_c_dim]
|
||||
c = observations[..., -self.latent_c_dim:]
|
||||
eps = self.style_activation(self.style_head(self.style_trunk(eps)))
|
||||
observations = torch.cat([obs, eps, c], dim=-1)
|
||||
embedding = self.actor_trunk(observations)
|
||||
return self.actor_head(embedding)
|
||||
274
Env/hbbc_background_policy.py
Normal file
274
Env/hbbc_background_policy.py
Normal file
@@ -0,0 +1,274 @@
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from Env.hbbc_actor_critic import ActorCritic
|
||||
|
||||
|
||||
def _wrap_to_pi(angle: float) -> float:
|
||||
return (angle + np.pi) % (2 * np.pi) - np.pi
|
||||
|
||||
|
||||
def _normalize_eps(eps: np.ndarray) -> np.ndarray:
|
||||
eps = np.asarray(eps, dtype=np.float32).reshape(-1)
|
||||
if eps.shape[0] != 6:
|
||||
raise ValueError(f"latent_eps must be 6-dim, got {eps.shape[0]}")
|
||||
norm = float(np.linalg.norm(eps))
|
||||
if norm < 1e-8:
|
||||
eps = np.array([1.0, 0.0, 0.0, 0.0, 0.0, 0.0], dtype=np.float32)
|
||||
else:
|
||||
eps = eps / norm
|
||||
return np.clip(eps, -1.0, 1.0)
|
||||
|
||||
|
||||
def _normalize_c(latent_c: np.ndarray) -> np.ndarray:
|
||||
c = np.asarray(latent_c, dtype=np.float32).reshape(-1)
|
||||
if c.shape[0] != 4:
|
||||
raise ValueError(f"latent_c must be 4-dim, got {c.shape[0]}")
|
||||
idx = int(np.argmax(c))
|
||||
one_hot = np.zeros(4, dtype=np.float32)
|
||||
one_hot[idx] = 1.0
|
||||
return one_hot
|
||||
|
||||
|
||||
def _sample_latent(rng: np.random.RandomState) -> Tuple[np.ndarray, np.ndarray]:
|
||||
eps = _normalize_eps(rng.randn(6).astype(np.float32))
|
||||
mode = int(rng.randint(0, 4))
|
||||
c = np.zeros(4, dtype=np.float32)
|
||||
c[mode] = 1.0
|
||||
return eps, c
|
||||
|
||||
|
||||
@dataclass
|
||||
class VehicleStateCache:
|
||||
last_heading_theta: Optional[float] = None
|
||||
last_action: Tuple[float, float] = (0.0, 0.0)
|
||||
last_speed_km_h: Optional[float] = None
|
||||
|
||||
|
||||
class HBBCModelWrapper:
|
||||
_cache: Dict[Tuple[str, str], "HBBCModelWrapper"] = {}
|
||||
|
||||
def __init__(self, model_path: str, device: str = "cpu"):
|
||||
self.model_path = os.path.abspath(model_path)
|
||||
self.device = torch.device(device)
|
||||
self.model = self._load_model()
|
||||
|
||||
@classmethod
|
||||
def get(cls, model_path: str, device: str = "cpu") -> "HBBCModelWrapper":
|
||||
key = (os.path.abspath(model_path), str(torch.device(device)))
|
||||
if key not in cls._cache:
|
||||
cls._cache[key] = HBBCModelWrapper(model_path=key[0], device=key[1])
|
||||
return cls._cache[key]
|
||||
|
||||
def _load_model(self) -> ActorCritic:
|
||||
model = ActorCritic(
|
||||
num_actor_obs=18,
|
||||
num_critic_obs=18,
|
||||
num_actions=2,
|
||||
latent_c_dim=4,
|
||||
latent_eps_dim=6,
|
||||
use_style_latent=True,
|
||||
).to(self.device)
|
||||
try:
|
||||
ckpt = torch.load(self.model_path, map_location=self.device, weights_only=True)
|
||||
except Exception:
|
||||
ckpt = torch.load(self.model_path, map_location=self.device, weights_only=False)
|
||||
state_dict = ckpt["actor_critic"] if isinstance(ckpt, dict) and "actor_critic" in ckpt else ckpt
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=False)
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
f"HBBC checkpoint missing required keys for {self.model_path}: {missing}"
|
||||
)
|
||||
if unexpected:
|
||||
print(f"[HBBC] ignore extra checkpoint keys: {unexpected[:8]}{'...' if len(unexpected) > 8 else ''}")
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
def act_batch(self, obs_batch: np.ndarray) -> np.ndarray:
|
||||
obs_batch = np.asarray(obs_batch, dtype=np.float32)
|
||||
with torch.no_grad():
|
||||
obs_t = torch.from_numpy(obs_batch).to(self.device)
|
||||
actions = self.model.act_inference(obs_t).cpu().numpy()
|
||||
return np.clip(actions, -1.0, 1.0)
|
||||
|
||||
|
||||
class HBBCLatentManager:
|
||||
def __init__(self, mode: str = "per_vehicle_fixed", seed: int = 0, latent_json_path: Optional[str] = None):
|
||||
self.mode = mode
|
||||
self.rng = np.random.RandomState(seed)
|
||||
self.latent_json_path = latent_json_path
|
||||
self.manual_object_latent: Dict[str, Dict[str, np.ndarray]] = {}
|
||||
self.manual_agent_latent: Dict[str, Dict[str, np.ndarray]] = {}
|
||||
self.manual_global_latent: Optional[Tuple[np.ndarray, np.ndarray]] = None
|
||||
self.vehicle_latent: Dict[str, Tuple[np.ndarray, np.ndarray]] = {}
|
||||
self._episode_latent: Optional[Tuple[np.ndarray, np.ndarray]] = None
|
||||
self._load_manual_latent_json()
|
||||
|
||||
def reset_episode(self):
|
||||
self.vehicle_latent.clear()
|
||||
self._episode_latent = None
|
||||
if self.mode == "per_episode_reset":
|
||||
self._episode_latent = _sample_latent(self.rng)
|
||||
|
||||
def _load_manual_latent_json(self):
|
||||
if not self.latent_json_path:
|
||||
return
|
||||
path = os.path.abspath(self.latent_json_path)
|
||||
if not os.path.exists(path):
|
||||
print(f"[HBBC] latent json not found: {path}, fallback to random sampling.")
|
||||
return
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
except Exception as e:
|
||||
print(f"[HBBC] failed to load latent json ({path}): {e}. fallback to random sampling.")
|
||||
return
|
||||
|
||||
object_section = data.get("object_id", {})
|
||||
agent_section = data.get("agent_id", {})
|
||||
global_section = data.get("global")
|
||||
|
||||
if global_section is not None:
|
||||
parsed = self._parse_one_latent(global_section, "global")
|
||||
if parsed is not None:
|
||||
self.manual_global_latent = (parsed["latent_eps"], parsed["latent_c"])
|
||||
|
||||
for key, value in object_section.items():
|
||||
parsed = self._parse_one_latent(value, f"object_id:{key}")
|
||||
if parsed is not None:
|
||||
self.manual_object_latent[str(key)] = parsed
|
||||
for key, value in agent_section.items():
|
||||
parsed = self._parse_one_latent(value, f"agent_id:{key}")
|
||||
if parsed is not None:
|
||||
self.manual_agent_latent[str(key)] = parsed
|
||||
|
||||
@staticmethod
|
||||
def _parse_one_latent(value: dict, name: str) -> Optional[Dict[str, np.ndarray]]:
|
||||
if not isinstance(value, dict):
|
||||
print(f"[HBBC] invalid latent entry ({name}): expect dict.")
|
||||
return None
|
||||
try:
|
||||
eps = _normalize_eps(value["latent_eps"])
|
||||
c = _normalize_c(value["latent_c"])
|
||||
return {"latent_eps": eps, "latent_c": c}
|
||||
except Exception as e:
|
||||
print(f"[HBBC] invalid latent entry ({name}): {e}")
|
||||
return None
|
||||
|
||||
def _lookup_manual(self, object_id: Optional[str], agent_id: Optional[str]) -> Optional[Tuple[np.ndarray, np.ndarray]]:
|
||||
if object_id is not None and object_id in self.manual_object_latent:
|
||||
e = self.manual_object_latent[object_id]["latent_eps"]
|
||||
c = self.manual_object_latent[object_id]["latent_c"]
|
||||
return e, c
|
||||
if agent_id is not None and agent_id in self.manual_agent_latent:
|
||||
e = self.manual_agent_latent[agent_id]["latent_eps"]
|
||||
c = self.manual_agent_latent[agent_id]["latent_c"]
|
||||
return e, c
|
||||
if self.manual_global_latent is not None:
|
||||
return self.manual_global_latent
|
||||
return None
|
||||
|
||||
def get_latent(self, vehicle_key: str, object_id: Optional[str], agent_id: Optional[str]) -> Tuple[np.ndarray, np.ndarray]:
|
||||
manual = self._lookup_manual(object_id=object_id, agent_id=agent_id)
|
||||
if manual is not None:
|
||||
return manual
|
||||
if self.mode == "per_episode_reset":
|
||||
if self._episode_latent is None:
|
||||
self._episode_latent = _sample_latent(self.rng)
|
||||
return self._episode_latent
|
||||
if vehicle_key not in self.vehicle_latent:
|
||||
self.vehicle_latent[vehicle_key] = _sample_latent(self.rng)
|
||||
return self.vehicle_latent[vehicle_key]
|
||||
|
||||
|
||||
class HBBCBackgroundController:
|
||||
def __init__(
|
||||
self,
|
||||
model_path: str,
|
||||
device: str = "cpu",
|
||||
latent_mode: str = "per_vehicle_fixed",
|
||||
latent_json_path: Optional[str] = None,
|
||||
seed: int = 0,
|
||||
dt: float = 0.1,
|
||||
):
|
||||
self.model = HBBCModelWrapper.get(model_path=model_path, device=device)
|
||||
self.latent_mgr = HBBCLatentManager(mode=latent_mode, seed=seed, latent_json_path=latent_json_path)
|
||||
self.dt = float(dt)
|
||||
self.vehicle_state: Dict[str, VehicleStateCache] = {}
|
||||
|
||||
def reset_episode(self):
|
||||
self.latent_mgr.reset_episode()
|
||||
self.vehicle_state.clear()
|
||||
|
||||
def remove_vehicle(self, vehicle_key: str):
|
||||
self.vehicle_state.pop(vehicle_key, None)
|
||||
self.latent_mgr.vehicle_latent.pop(vehicle_key, None)
|
||||
|
||||
def _build_base_state(self, vehicle, vehicle_key: str) -> np.ndarray:
|
||||
state = self.vehicle_state.get(vehicle_key)
|
||||
if state is None:
|
||||
state = VehicleStateCache()
|
||||
self.vehicle_state[vehicle_key] = state
|
||||
|
||||
speed_km_h = float(getattr(vehicle, "speed_km_h", 0.0))
|
||||
max_speed_km_h = float(getattr(vehicle, "max_speed_km_h", 120.0))
|
||||
veh_vel = np.clip((speed_km_h + 1.0) / (max_speed_km_h + 1.0), 0.0, 1.0)
|
||||
|
||||
heading_theta = float(getattr(vehicle, "heading_theta", 0.0))
|
||||
if state.last_heading_theta is None:
|
||||
yaw_rate = 0.0
|
||||
else:
|
||||
yaw_rate = _wrap_to_pi(heading_theta - state.last_heading_theta) / self.dt
|
||||
yaw_rate = float(np.clip(yaw_rate, -5.0, 5.0))
|
||||
|
||||
current_action = getattr(vehicle, "current_action", None)
|
||||
if current_action is None:
|
||||
last_action_0, last_action_1 = state.last_action
|
||||
else:
|
||||
try:
|
||||
last_action_0, last_action_1 = float(current_action[0]), float(current_action[1])
|
||||
except Exception:
|
||||
last_action_0, last_action_1 = state.last_action
|
||||
|
||||
state.last_heading_theta = heading_theta
|
||||
state.last_speed_km_h = speed_km_h
|
||||
state.last_action = (last_action_0, last_action_1)
|
||||
|
||||
obs = np.array(
|
||||
[
|
||||
0.0,
|
||||
0.0,
|
||||
0.0,
|
||||
veh_vel,
|
||||
0.0,
|
||||
yaw_rate * 0.5,
|
||||
last_action_0,
|
||||
last_action_1,
|
||||
],
|
||||
dtype=np.float32,
|
||||
)
|
||||
return obs
|
||||
|
||||
def build_obs(self, vehicle, vehicle_key: str, object_id: Optional[str], agent_id: Optional[str]) -> np.ndarray:
|
||||
base = self._build_base_state(vehicle, vehicle_key=vehicle_key)
|
||||
eps, c = self.latent_mgr.get_latent(vehicle_key=vehicle_key, object_id=object_id, agent_id=agent_id)
|
||||
return np.concatenate([base, eps, c], axis=-1).astype(np.float32)
|
||||
|
||||
def infer_actions(self, batch: List[Tuple[str, object, Optional[str], Optional[str]]]) -> Dict[str, np.ndarray]:
|
||||
if not batch:
|
||||
return {}
|
||||
obs_list = []
|
||||
vehicle_ids = []
|
||||
for vehicle_key, vehicle, object_id, agent_id in batch:
|
||||
obs_list.append(self.build_obs(vehicle, vehicle_key=vehicle_key, object_id=object_id, agent_id=agent_id))
|
||||
vehicle_ids.append(vehicle_key)
|
||||
actions = self.model.act_batch(np.stack(obs_list, axis=0))
|
||||
out = {}
|
||||
for idx, key in enumerate(vehicle_ids):
|
||||
out[key] = actions[idx].astype(np.float32)
|
||||
return out
|
||||
93
Env/inverse_dynamics.py
Normal file
93
Env/inverse_dynamics.py
Normal file
@@ -0,0 +1,93 @@
|
||||
import numpy as np
|
||||
import math
|
||||
|
||||
class InverseDynamics:
|
||||
def __init__(self, max_steering=0.7, max_acc=8.0, length=4.5):
|
||||
"""
|
||||
:param max_steering: Max steering angle in radians (approx 40 degrees)
|
||||
:param max_acc: Max acceleration in m/s^2
|
||||
:param length: Vehicle length in meters (Waymo default approx 4.5m)
|
||||
"""
|
||||
self.max_steering = max_steering
|
||||
self.max_acc = max_acc
|
||||
self.wheelbase = 0.7 * length # Approximation as per request
|
||||
|
||||
def compute_action(self, current_state, next_state, dt=0.1):
|
||||
"""
|
||||
Compute action [steering, acceleration] from current and next state.
|
||||
State format: dictionary or object with keys/attrs: position (x, y), heading, velocity (v_x, v_y)
|
||||
or numpy array [x, y, vx, vy, heading]
|
||||
|
||||
Using Bicycle Model:
|
||||
delta = arctan(L * theta_dot / v)
|
||||
acc = (v_next - v_curr) / dt
|
||||
"""
|
||||
|
||||
# Extract state
|
||||
# Assume state is dict-like for now, can adapt if needed
|
||||
# We need: velocity (scalar), heading
|
||||
|
||||
# Helper to get speed
|
||||
def get_speed(vel):
|
||||
return np.linalg.norm(vel)
|
||||
|
||||
v_curr = get_speed(current_state['velocity'])
|
||||
v_next = get_speed(next_state['velocity'])
|
||||
|
||||
# 1. Acceleration (longitudinal)
|
||||
acc = (v_next - v_curr) / dt
|
||||
|
||||
# 2. Steering (lateral)
|
||||
# theta_dot = (theta_next - theta_curr) / dt
|
||||
theta_curr = current_state['heading']
|
||||
theta_next = next_state['heading']
|
||||
|
||||
# Handle angle wrapping [-pi, pi]
|
||||
diff_theta = theta_next - theta_curr
|
||||
if diff_theta > np.pi:
|
||||
diff_theta -= 2 * np.pi
|
||||
elif diff_theta < -np.pi:
|
||||
diff_theta += 2 * np.pi
|
||||
|
||||
theta_dot = diff_theta / dt
|
||||
|
||||
# Avoid division by zero for stationary vehicles
|
||||
if v_curr < 0.1:
|
||||
steering = 0.0
|
||||
else:
|
||||
# delta = arctan(L * theta_dot / v)
|
||||
steering = np.arctan(self.wheelbase * theta_dot / v_curr)
|
||||
|
||||
# Normalize actions to [-1, 1]
|
||||
norm_acc = np.clip(acc / self.max_acc, -1.0, 1.0)
|
||||
norm_steering = np.clip(steering / self.max_steering, -1.0, 1.0)
|
||||
|
||||
return np.array([norm_steering, norm_acc]), {'raw_acc': acc, 'raw_steering': steering}
|
||||
|
||||
def apply_action(self, current_state, action, dt=0.1):
|
||||
"""
|
||||
Forward dynamics: given current_state and action [steering, acc] in [-1, 1], return next_state.
|
||||
State format: dict with position (x,y), heading, velocity (vx, vy).
|
||||
"""
|
||||
steering_norm, acc_norm = float(action[0]), float(action[1])
|
||||
acc = acc_norm * self.max_acc
|
||||
steering = steering_norm * self.max_steering
|
||||
pos = np.array(current_state['position'][:2], dtype=np.float64)
|
||||
heading = float(current_state['heading'])
|
||||
vel = np.array(current_state['velocity'], dtype=np.float64)
|
||||
v = np.linalg.norm(vel)
|
||||
if v < 0.1:
|
||||
v = 0.1
|
||||
theta_dot = v * np.tan(steering) / self.wheelbase
|
||||
v_next = v + acc * dt
|
||||
v_next = max(0.0, v_next)
|
||||
heading_next = heading + theta_dot * dt
|
||||
heading_next = np.arctan2(np.sin(heading_next), np.cos(heading_next))
|
||||
vx_next = v_next * np.cos(heading_next)
|
||||
vy_next = v_next * np.sin(heading_next)
|
||||
pos_next = pos + dt * np.array([vx_next, vy_next])
|
||||
return {
|
||||
'position': pos_next,
|
||||
'heading': heading_next,
|
||||
'velocity': np.array([vx_next, vy_next]),
|
||||
}
|
||||
@@ -1,170 +0,0 @@
|
||||
"""
|
||||
日志工具模块
|
||||
提供将终端输出同时保存到文件的功能
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
class TeeLogger:
|
||||
"""
|
||||
双向输出类:同时输出到终端和文件
|
||||
"""
|
||||
def __init__(self, filename, mode='w', terminal=None):
|
||||
"""
|
||||
Args:
|
||||
filename: 日志文件路径
|
||||
mode: 文件打开模式 ('w'=覆盖, 'a'=追加)
|
||||
terminal: 原始输出流(通常是sys.stdout或sys.stderr)
|
||||
"""
|
||||
self.terminal = terminal or sys.stdout
|
||||
self.log_file = open(filename, mode, encoding='utf-8')
|
||||
|
||||
def write(self, message):
|
||||
"""写入消息到终端和文件"""
|
||||
self.terminal.write(message)
|
||||
self.log_file.write(message)
|
||||
self.log_file.flush() # 立即写入磁盘
|
||||
|
||||
def flush(self):
|
||||
"""刷新缓冲区"""
|
||||
self.terminal.flush()
|
||||
self.log_file.flush()
|
||||
|
||||
def close(self):
|
||||
"""关闭日志文件"""
|
||||
if self.log_file:
|
||||
self.log_file.close()
|
||||
|
||||
|
||||
class LoggerContext:
|
||||
"""
|
||||
日志上下文管理器
|
||||
使用with语句自动管理日志的开启和关闭
|
||||
"""
|
||||
def __init__(self, log_file=None, log_dir="logs", mode='w',
|
||||
redirect_stdout=True, redirect_stderr=True):
|
||||
"""
|
||||
Args:
|
||||
log_file: 日志文件名(None则自动生成时间戳文件名)
|
||||
log_dir: 日志目录
|
||||
mode: 文件打开模式 ('w'=覆盖, 'a'=追加)
|
||||
redirect_stdout: 是否重定向标准输出
|
||||
redirect_stderr: 是否重定向标准错误
|
||||
"""
|
||||
self.log_dir = log_dir
|
||||
self.mode = mode
|
||||
self.redirect_stdout = redirect_stdout
|
||||
self.redirect_stderr = redirect_stderr
|
||||
|
||||
# 创建日志目录
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
|
||||
# 生成日志文件名
|
||||
if log_file is None:
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
log_file = f"run_{timestamp}.log"
|
||||
|
||||
self.log_path = os.path.join(log_dir, log_file)
|
||||
|
||||
# 保存原始的stdout和stderr
|
||||
self.original_stdout = sys.stdout
|
||||
self.original_stderr = sys.stderr
|
||||
|
||||
# 日志对象
|
||||
self.stdout_logger = None
|
||||
self.stderr_logger = None
|
||||
|
||||
def __enter__(self):
|
||||
"""进入上下文:开启日志"""
|
||||
print(f"📝 日志记录已启用")
|
||||
print(f"📁 日志文件: {self.log_path}")
|
||||
print("-" * 60)
|
||||
|
||||
# 创建TeeLogger对象
|
||||
if self.redirect_stdout:
|
||||
self.stdout_logger = TeeLogger(
|
||||
self.log_path,
|
||||
mode=self.mode,
|
||||
terminal=self.original_stdout
|
||||
)
|
||||
sys.stdout = self.stdout_logger
|
||||
|
||||
if self.redirect_stderr:
|
||||
self.stderr_logger = TeeLogger(
|
||||
self.log_path,
|
||||
mode='a', # stderr总是追加模式
|
||||
terminal=self.original_stderr
|
||||
)
|
||||
sys.stderr = self.stderr_logger
|
||||
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
"""退出上下文:关闭日志"""
|
||||
# 恢复原始输出
|
||||
sys.stdout = self.original_stdout
|
||||
sys.stderr = self.original_stderr
|
||||
|
||||
# 关闭日志文件
|
||||
if self.stdout_logger:
|
||||
self.stdout_logger.close()
|
||||
if self.stderr_logger:
|
||||
self.stderr_logger.close()
|
||||
|
||||
print("-" * 60)
|
||||
print(f"✅ 日志已保存到: {self.log_path}")
|
||||
|
||||
# 返回False表示不抑制异常
|
||||
return False
|
||||
|
||||
|
||||
def setup_logger(log_file=None, log_dir="logs", mode='w'):
|
||||
"""
|
||||
快速设置日志记录
|
||||
|
||||
Args:
|
||||
log_file: 日志文件名(None则自动生成)
|
||||
log_dir: 日志目录
|
||||
mode: 文件模式 ('w'=覆盖, 'a'=追加)
|
||||
|
||||
Returns:
|
||||
LoggerContext对象
|
||||
|
||||
Example:
|
||||
with setup_logger("my_test.log"):
|
||||
print("这条消息会同时输出到终端和文件")
|
||||
"""
|
||||
return LoggerContext(log_file=log_file, log_dir=log_dir, mode=mode)
|
||||
|
||||
|
||||
def get_default_log_filename(prefix="run"):
|
||||
"""
|
||||
生成默认的日志文件名(带时间戳)
|
||||
|
||||
Args:
|
||||
prefix: 文件名前缀
|
||||
|
||||
Returns:
|
||||
str: 格式为 "prefix_YYYYMMDD_HHMMSS.log"
|
||||
"""
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
return f"{prefix}_{timestamp}.log"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 测试代码
|
||||
print("测试1: 使用默认配置")
|
||||
with setup_logger():
|
||||
print("这是测试消息1")
|
||||
print("这是测试消息2")
|
||||
print("日志记录已结束\n")
|
||||
|
||||
print("测试2: 使用自定义文件名")
|
||||
with setup_logger(log_file="test_custom.log"):
|
||||
print("自定义文件名测试")
|
||||
for i in range(3):
|
||||
print(f" 消息 {i+1}")
|
||||
print("完成")
|
||||
|
||||
@@ -1,20 +1,10 @@
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
from Env.simple_idm_policy import ConstantVelocityPolicy
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
from logger_utils import setup_logger
|
||||
import sys
|
||||
import os
|
||||
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/Env"
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/data"
|
||||
|
||||
def main(enable_logging=False, log_file=None):
|
||||
"""
|
||||
主函数
|
||||
|
||||
Args:
|
||||
enable_logging: 是否启用日志记录到文件
|
||||
log_file: 日志文件名(None则自动生成时间戳文件名)
|
||||
"""
|
||||
def main():
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
# "data_directory": AssetLoader.file_path(AssetLoader.asset_path, "waymo", unix_style=False),
|
||||
@@ -26,18 +16,12 @@ def main(enable_logging=False, log_file=None):
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": True,
|
||||
"manual_control": True,
|
||||
|
||||
# 车道检测与过滤配置
|
||||
"filter_offroad_vehicles": True, # 启用车道区域过滤,过滤草坪等非车道区域的车辆
|
||||
"lane_tolerance": 3.0, # 车道检测容差(米),可根据需要调整
|
||||
"max_controlled_vehicles": 2, # 限制最大车辆数(可选,None表示不限制)
|
||||
"debug_lane_filter": True,
|
||||
"debug_traffic_light": True,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
obs = env.reset(0)
|
||||
obs = env.reset(0
|
||||
)
|
||||
for step in range(10000):
|
||||
actions = {
|
||||
aid: env.controlled_agents[aid].policy.act()
|
||||
@@ -54,25 +38,4 @@ def main(enable_logging=False, log_file=None):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 解析命令行参数
|
||||
enable_logging = "--log" in sys.argv or "-l" in sys.argv
|
||||
|
||||
# 提取自定义日志文件名
|
||||
log_file = None
|
||||
for arg in sys.argv:
|
||||
if arg.startswith("--log-file="):
|
||||
log_file = arg.split("=")[1]
|
||||
break
|
||||
|
||||
if enable_logging:
|
||||
# 使用日志记录
|
||||
log_dir = os.path.join(os.path.dirname(__file__), "logs")
|
||||
with setup_logger(log_file=log_file, log_dir=log_dir):
|
||||
main(enable_logging=True, log_file=log_file)
|
||||
else:
|
||||
# 普通运行(只输出到终端)
|
||||
print("💡 提示: 使用 --log 或 -l 参数启用日志记录")
|
||||
print(" 示例: python run_multiagent_env.py --log")
|
||||
print(" 自定义文件名: python run_multiagent_env.py --log --log-file=my_run.log")
|
||||
print("-" * 60)
|
||||
main(enable_logging=False)
|
||||
main()
|
||||
@@ -1,115 +0,0 @@
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
from logger_utils import setup_logger
|
||||
import time
|
||||
import sys
|
||||
import os
|
||||
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/Env"
|
||||
|
||||
def main(enable_logging=False):
|
||||
"""极致性能优化版本 - 启用所有优化选项"""
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 300,
|
||||
|
||||
# 关闭所有渲染
|
||||
"use_render": False,
|
||||
"render_pipeline": False,
|
||||
"image_observation": False,
|
||||
"interface_panel": [],
|
||||
"manual_control": False,
|
||||
"show_fps": False,
|
||||
"debug": False,
|
||||
|
||||
# 物理引擎优化
|
||||
"physics_world_step_size": 0.02,
|
||||
"decision_repeat": 5,
|
||||
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": True,
|
||||
|
||||
# 车道检测与过滤配置
|
||||
"filter_offroad_vehicles": True, # 过滤非车道区域的车辆
|
||||
"lane_tolerance": 3.0,
|
||||
"max_controlled_vehicles": 15, # 限制车辆数以提升性能
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
# 【关键优化】启用激光雷达缓存
|
||||
# 每3帧才重新计算激光雷达,其余帧使用缓存
|
||||
# 可将激光雷达计算量减少到原来的1/3
|
||||
env.lidar_cache_interval = 3
|
||||
|
||||
obs = env.reset(0)
|
||||
|
||||
# 性能统计
|
||||
start_time = time.time()
|
||||
total_steps = 0
|
||||
|
||||
print("=" * 60)
|
||||
print("极致性能模式")
|
||||
print("激光雷达优化:80→40束 (前向), 10→6束 (侧向+车道线)")
|
||||
print("激光雷达缓存:每3帧计算一次,中间帧使用缓存")
|
||||
print("预期性能提升:3-5倍")
|
||||
print("=" * 60)
|
||||
|
||||
for step in range(10000):
|
||||
actions = {
|
||||
aid: env.controlled_agents[aid].policy.act()
|
||||
for aid in env.controlled_agents
|
||||
}
|
||||
|
||||
obs, rewards, dones, infos = env.step(actions)
|
||||
total_steps += 1
|
||||
|
||||
# 每100步输出一次性能统计
|
||||
if step % 100 == 0 and step > 0:
|
||||
elapsed = time.time() - start_time
|
||||
fps = total_steps / elapsed
|
||||
print(f"Step {step:4d}: FPS = {fps:6.2f}, 车辆数 = {len(env.controlled_agents):3d}, "
|
||||
f"平均步时间 = {1000/fps:.2f}ms")
|
||||
|
||||
if dones["__all__"]:
|
||||
break
|
||||
|
||||
# 最终统计
|
||||
elapsed = time.time() - start_time
|
||||
fps = total_steps / elapsed
|
||||
print("\n" + "=" * 60)
|
||||
print(f"总计: {total_steps} 步")
|
||||
print(f"耗时: {elapsed:.2f}s")
|
||||
print(f"平均FPS: {fps:.2f}")
|
||||
print(f"单步平均耗时: {1000/fps:.2f}ms")
|
||||
print("=" * 60)
|
||||
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 解析命令行参数
|
||||
enable_logging = "--log" in sys.argv or "-l" in sys.argv
|
||||
|
||||
# 提取自定义日志文件名
|
||||
log_file = None
|
||||
for arg in sys.argv:
|
||||
if arg.startswith("--log-file="):
|
||||
log_file = arg.split("=")[1]
|
||||
break
|
||||
|
||||
if enable_logging:
|
||||
# 使用日志记录
|
||||
log_dir = os.path.join(os.path.dirname(__file__), "logs")
|
||||
with setup_logger(log_file=log_file or "run_fast.log", log_dir=log_dir):
|
||||
main(enable_logging=True)
|
||||
else:
|
||||
# 普通运行(只输出到终端)
|
||||
print("💡 提示: 使用 --log 或 -l 参数启用日志记录")
|
||||
print("-" * 60)
|
||||
main(enable_logging=False)
|
||||
|
||||
@@ -1,156 +0,0 @@
|
||||
"""
|
||||
多进程并行版本 - 充分利用多核CPU
|
||||
适合大规模数据收集和训练
|
||||
"""
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
import time
|
||||
import os
|
||||
from multiprocessing import Pool, cpu_count
|
||||
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/Env"
|
||||
|
||||
|
||||
def run_single_env(args):
|
||||
"""在单个进程中运行一个环境实例"""
|
||||
seed, num_steps, worker_id = args
|
||||
|
||||
# 创建环境(每个进程独立)
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 300,
|
||||
|
||||
# 性能优化
|
||||
"use_render": False,
|
||||
"render_pipeline": False,
|
||||
"image_observation": False,
|
||||
"interface_panel": [],
|
||||
"manual_control": False,
|
||||
"show_fps": False,
|
||||
"debug": False,
|
||||
|
||||
"physics_world_step_size": 0.02,
|
||||
"decision_repeat": 5,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": True,
|
||||
|
||||
# 车道检测与过滤配置
|
||||
"filter_offroad_vehicles": True,
|
||||
"lane_tolerance": 3.0,
|
||||
"max_controlled_vehicles": 15,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
# 启用激光雷达缓存
|
||||
env.lidar_cache_interval = 3
|
||||
|
||||
# 运行仿真
|
||||
start_time = time.time()
|
||||
obs = env.reset(seed)
|
||||
total_steps = 0
|
||||
total_agents = 0
|
||||
|
||||
for step in range(num_steps):
|
||||
actions = {
|
||||
aid: env.controlled_agents[aid].policy.act()
|
||||
for aid in env.controlled_agents
|
||||
}
|
||||
|
||||
obs, rewards, dones, infos = env.step(actions)
|
||||
total_steps += 1
|
||||
total_agents += len(env.controlled_agents)
|
||||
|
||||
if dones["__all__"]:
|
||||
break
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
fps = total_steps / elapsed if elapsed > 0 else 0
|
||||
avg_agents = total_agents / total_steps if total_steps > 0 else 0
|
||||
|
||||
env.close()
|
||||
|
||||
return {
|
||||
'worker_id': worker_id,
|
||||
'seed': seed,
|
||||
'steps': total_steps,
|
||||
'elapsed': elapsed,
|
||||
'fps': fps,
|
||||
'avg_agents': avg_agents,
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
"""主函数:协调多个并行环境"""
|
||||
# 获取CPU核心数
|
||||
num_cores = cpu_count()
|
||||
# 建议使用物理核心数(12600KF是10核20线程,使用10个进程)
|
||||
num_workers = min(10, num_cores)
|
||||
|
||||
print("=" * 80)
|
||||
print(f"多进程并行模式")
|
||||
print(f"CPU核心数: {num_cores}")
|
||||
print(f"并行进程数: {num_workers}")
|
||||
print(f"每个环境运行: 1000步")
|
||||
print("=" * 80)
|
||||
|
||||
# 准备任务参数
|
||||
num_steps_per_env = 1000
|
||||
tasks = [(seed, num_steps_per_env, worker_id)
|
||||
for worker_id, seed in enumerate(range(num_workers))]
|
||||
|
||||
# 启动多进程池
|
||||
start_time = time.time()
|
||||
|
||||
with Pool(processes=num_workers) as pool:
|
||||
results = pool.map(run_single_env, tasks)
|
||||
|
||||
total_elapsed = time.time() - start_time
|
||||
|
||||
# 统计结果
|
||||
print("\n" + "=" * 80)
|
||||
print("各进程执行结果:")
|
||||
print("-" * 80)
|
||||
print(f"{'Worker':<8} {'Seed':<6} {'Steps':<8} {'Time(s)':<10} {'FPS':<8} {'平均车辆数':<12}")
|
||||
print("-" * 80)
|
||||
|
||||
total_steps = 0
|
||||
total_fps = 0
|
||||
|
||||
for result in results:
|
||||
print(f"{result['worker_id']:<8} "
|
||||
f"{result['seed']:<6} "
|
||||
f"{result['steps']:<8} "
|
||||
f"{result['elapsed']:<10.2f} "
|
||||
f"{result['fps']:<8.2f} "
|
||||
f"{result['avg_agents']:<12.1f}")
|
||||
total_steps += result['steps']
|
||||
total_fps += result['fps']
|
||||
|
||||
print("-" * 80)
|
||||
avg_fps_per_env = total_fps / len(results)
|
||||
total_throughput = total_steps / total_elapsed
|
||||
|
||||
print(f"\n总体统计:")
|
||||
print(f" 总步数: {total_steps}")
|
||||
print(f" 总耗时: {total_elapsed:.2f}s")
|
||||
print(f" 单环境平均FPS: {avg_fps_per_env:.2f}")
|
||||
print(f" 总吞吐量: {total_throughput:.2f} steps/s")
|
||||
print(f" 并行效率: {total_throughput / avg_fps_per_env:.1f}x")
|
||||
print("=" * 80)
|
||||
|
||||
# 与单进程对比
|
||||
print(f"\n性能对比:")
|
||||
print(f" 单进程FPS (预估): ~30 FPS")
|
||||
print(f" 多进程吞吐量: {total_throughput:.2f} steps/s")
|
||||
print(f" 性能提升: {total_throughput / 30:.1f}x")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
import time
|
||||
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/Env"
|
||||
|
||||
def main():
|
||||
"""带可视化的版本(低FPS,约15帧)"""
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 300,
|
||||
|
||||
# 可视化设置(牺牲性能)
|
||||
"use_render": True,
|
||||
"manual_control": False,
|
||||
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": True,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
obs = env.reset(0)
|
||||
|
||||
start_time = time.time()
|
||||
total_steps = 0
|
||||
|
||||
for step in range(10000):
|
||||
actions = {
|
||||
aid: env.controlled_agents[aid].policy.act()
|
||||
for aid in env.controlled_agents
|
||||
}
|
||||
|
||||
obs, rewards, dones, infos = env.step(actions)
|
||||
env.render(mode="topdown") # 实时渲染
|
||||
|
||||
total_steps += 1
|
||||
|
||||
if step % 100 == 0 and step > 0:
|
||||
elapsed = time.time() - start_time
|
||||
fps = total_steps / elapsed
|
||||
print(f"Step {step}: FPS = {fps:.2f}, 车辆数 = {len(env.controlled_agents)}")
|
||||
|
||||
if dones["__all__"]:
|
||||
break
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
fps = total_steps / elapsed
|
||||
print(f"\n总计: {total_steps} 步,耗时 {elapsed:.2f}s,平均FPS = {fps:.2f}")
|
||||
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -53,13 +53,13 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
data_directory=None,
|
||||
num_controlled_agents=3,
|
||||
horizon=1000,
|
||||
# 车道检测与过滤配置
|
||||
filter_offroad_vehicles=True, # 是否过滤非车道区域的车辆
|
||||
lane_tolerance=3.0, # 车道检测容差(米),用于放宽边界条件
|
||||
max_controlled_vehicles=None, # 最大可控车辆数限制(None表示不限制)
|
||||
# 调试模式配置
|
||||
debug_traffic_light=False, # 是否启用红绿灯检测调试输出
|
||||
debug_lane_filter=False, # 是否启用车道过滤调试输出
|
||||
# HBBC background vehicle controls (optional)
|
||||
enable_hbbc_background=False,
|
||||
hbbc_model_path="models/hbbc/hbbc.pt",
|
||||
hbbc_inference_device="cpu",
|
||||
hbbc_latent_mode="per_vehicle_fixed",
|
||||
hbbc_latent_json_path=None,
|
||||
hbbc_dt=0.1,
|
||||
))
|
||||
return config
|
||||
|
||||
@@ -69,11 +69,13 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
self.controlled_agent_ids = []
|
||||
self.obs_list = []
|
||||
self.round = 0
|
||||
# 调试模式配置
|
||||
self.debug_traffic_light = config.get("debug_traffic_light", False)
|
||||
self.debug_lane_filter = config.get("debug_lane_filter", False)
|
||||
super().__init__(config)
|
||||
|
||||
@property
|
||||
def num_controlled_in_scenario(self) -> int:
|
||||
"""整个场景中受控车轨迹总数(car_birth_info_list 长度),会在不同 show_time 陆续 spawn。"""
|
||||
return len(getattr(self, "car_birth_info_list", []))
|
||||
|
||||
def reset(self, seed: Union[None, int] = None):
|
||||
self.round = 0
|
||||
if self.logger is None:
|
||||
@@ -86,76 +88,28 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
if self.engine is None:
|
||||
raise ValueError("Broken MetaDrive instance.")
|
||||
|
||||
# 记录专家数据中每辆车的位置,接着全部清除,只保留位置等信息,用于后续生成
|
||||
_obj_to_clean_this_frame = []
|
||||
self.car_birth_info_list = []
|
||||
for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items():
|
||||
if scenario_id == self.engine.traffic_manager.sdc_scenario_id:
|
||||
continue
|
||||
else:
|
||||
if track["type"] == MetaDriveType.VEHICLE:
|
||||
_obj_to_clean_this_frame.append(scenario_id)
|
||||
valid = track['state']['valid']
|
||||
first_show = np.argmax(valid) if valid.any() else -1
|
||||
last_show = len(valid) - 1 - np.argmax(valid[::-1]) if valid.any() else -1
|
||||
# id,出现时间,出生点坐标,出生朝向,目的地
|
||||
self.car_birth_info_list.append({
|
||||
'id': track['metadata']['object_id'],
|
||||
'show_time': first_show,
|
||||
'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]),
|
||||
'heading': track['state']['heading'][first_show],
|
||||
'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1])
|
||||
})
|
||||
|
||||
for scenario_id in _obj_to_clean_this_frame:
|
||||
# 注意:_build_birth_lists_from_traffic() 在 engine.reset() 之前执行,读的是当前 engine 的
|
||||
# current_traffic_data 与 map_manager.current_map。若复用同一 env 连续 reset(0)、reset(1),
|
||||
# MetaDrive 可能已按 seed 更新了 traffic 为 scenario 1,但 map 仍为 scenario 0(在 engine.reset() 才切图),
|
||||
# 导致 is_on_lane( scenario_1 车位, scenario_0 地图 ) 全为 False → 全部 off_lane → 0 受控车。
|
||||
# 因此多场景时应“每个 scenario 单独建 env”(start_scenario_index=i, num_scenarios=1)再 reset(seed=i)。
|
||||
self.background_vehicles = getattr(self, "background_vehicles", {})
|
||||
self.car_birth_info_list, self.background_vehicles, _obj_to_clean = self._build_birth_lists_from_traffic()
|
||||
for scenario_id in _obj_to_clean:
|
||||
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()
|
||||
|
||||
self.lanes = self.engine.map_manager.current_map.road_network.graph
|
||||
|
||||
# 调试:场景信息统计
|
||||
if self.debug_lane_filter or self.debug_traffic_light:
|
||||
print(f"\n📍 场景信息统计:")
|
||||
print(f" - 总车道数: {len(self.lanes)}")
|
||||
|
||||
# 统计红绿灯数量
|
||||
if self.debug_traffic_light:
|
||||
traffic_light_lanes = []
|
||||
for lane in self.lanes.values():
|
||||
if self.engine.light_manager.has_traffic_light(lane.lane.index):
|
||||
traffic_light_lanes.append(lane.lane.index)
|
||||
print(f" - 有红绿灯的车道数: {len(traffic_light_lanes)}")
|
||||
if len(traffic_light_lanes) > 0:
|
||||
print(f" 车道索引: {traffic_light_lanes[:5]}" +
|
||||
(f" ... 共{len(traffic_light_lanes)}个" if len(traffic_light_lanes) > 5 else ""))
|
||||
else:
|
||||
print(f" ⚠️ 场景中没有红绿灯!")
|
||||
|
||||
# 在获取车道信息后,进行车道区域过滤
|
||||
total_cars_before = len(self.car_birth_info_list)
|
||||
valid_count, filtered_count, filtered_list = self._filter_valid_spawn_positions()
|
||||
|
||||
# 输出过滤信息
|
||||
if filtered_count > 0:
|
||||
self.logger.warning(f"车辆生成位置过滤: 原始 {total_cars_before} 辆, "
|
||||
f"有效 {valid_count} 辆, 过滤 {filtered_count} 辆")
|
||||
for filtered_car in filtered_list[:5]: # 只显示前5个
|
||||
self.logger.debug(f" - 过滤车辆 ID={filtered_car['id']}, "
|
||||
f"位置={filtered_car['position']}, "
|
||||
f"原因={filtered_car['reason']}")
|
||||
if filtered_count > 5:
|
||||
self.logger.debug(f" - ... 还有 {filtered_count - 5} 辆车被过滤")
|
||||
|
||||
# 限制最大车辆数(在过滤后应用)
|
||||
max_vehicles = self.config.get("max_controlled_vehicles", None)
|
||||
if max_vehicles is not None and len(self.car_birth_info_list) > max_vehicles:
|
||||
self.car_birth_info_list = self.car_birth_info_list[:max_vehicles]
|
||||
self.logger.info(f"限制最大车辆数为 {max_vehicles} 辆")
|
||||
|
||||
self.logger.info(f"最终生成 {len(self.car_birth_info_list)} 辆可控车辆")
|
||||
|
||||
if self.top_down_renderer is not None:
|
||||
self.top_down_renderer.clear()
|
||||
@@ -165,116 +119,32 @@ 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()
|
||||
|
||||
return self._get_all_obs()
|
||||
|
||||
def _is_position_on_lane(self, position, tolerance=None):
|
||||
"""
|
||||
检测给定位置是否在有效车道范围内
|
||||
|
||||
Args:
|
||||
position: (x, y) 车辆位置坐标
|
||||
tolerance: 容差范围(米),用于放宽检测条件。None时使用配置中的默认值
|
||||
|
||||
Returns:
|
||||
bool: True表示在车道上,False表示在非车道区域(如草坪、停车场等)
|
||||
"""
|
||||
if not hasattr(self, 'lanes') or self.lanes is None:
|
||||
if self.debug_lane_filter:
|
||||
print(f" ⚠️ 车道信息未初始化,默认允许")
|
||||
return True # 如果车道信息未初始化,默认允许生成
|
||||
|
||||
if tolerance is None:
|
||||
tolerance = self.config.get("lane_tolerance", 3.0)
|
||||
|
||||
position_2d = (position[0], position[1])
|
||||
|
||||
if self.debug_lane_filter:
|
||||
print(f" 🔍 检测位置 ({position_2d[0]:.2f}, {position_2d[1]:.2f}), 容差={tolerance}m")
|
||||
|
||||
# 方法1:直接检测是否在任一车道上
|
||||
checked_lanes = 0
|
||||
for lane in self.lanes.values():
|
||||
try:
|
||||
checked_lanes += 1
|
||||
if lane.lane.point_on_lane(position_2d):
|
||||
if self.debug_lane_filter:
|
||||
print(f" ✅ 在车道上 (车道{lane.lane.index}, 检查了{checked_lanes}条)")
|
||||
return True
|
||||
except:
|
||||
def _build_birth_lists_from_traffic(self):
|
||||
"""Build car_birth_info_list and obj_to_clean from current_traffic_data. Override for filtered (lane/static) selection."""
|
||||
_obj_to_clean_this_frame = []
|
||||
car_birth_info_list = []
|
||||
for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items():
|
||||
if scenario_id == self.engine.traffic_manager.sdc_scenario_id:
|
||||
continue
|
||||
|
||||
if self.debug_lane_filter:
|
||||
print(f" ❌ 不在任何车道上 (检查了{checked_lanes}条车道)")
|
||||
|
||||
# 方法2:如果严格检测失败,使用容差范围检测(考虑车道边缘)
|
||||
# 注释:此方法已被禁用,如需启用请取消注释
|
||||
# if tolerance > 0:
|
||||
# for lane in self.lanes.values():
|
||||
# try:
|
||||
# # 计算点到车道中心线的距离
|
||||
# lane_obj = lane.lane
|
||||
# # 获取车道长度并检测最近点
|
||||
# s, lateral = lane_obj.local_coordinates(position_2d)
|
||||
|
||||
# # 如果横向距离在容差范围内,认为是有效的
|
||||
# if abs(lateral) <= tolerance and 0 <= s <= lane_obj.length:
|
||||
# return True
|
||||
# except:
|
||||
# continue
|
||||
|
||||
return False
|
||||
|
||||
def _filter_valid_spawn_positions(self):
|
||||
"""
|
||||
过滤掉生成位置不在有效车道上的车辆信息
|
||||
根据配置决定是否执行过滤
|
||||
|
||||
Returns:
|
||||
tuple: (有效车辆数量, 被过滤车辆数量, 被过滤车辆ID列表)
|
||||
"""
|
||||
# 如果配置中禁用了过滤,直接返回
|
||||
if not self.config.get("filter_offroad_vehicles", True):
|
||||
if self.debug_lane_filter:
|
||||
print(f"🚫 车道过滤已禁用")
|
||||
return len(self.car_birth_info_list), 0, []
|
||||
|
||||
if self.debug_lane_filter:
|
||||
print(f"\n🔍 开始车道过滤: 共 {len(self.car_birth_info_list)} 辆车待检测")
|
||||
|
||||
valid_cars = []
|
||||
filtered_cars = []
|
||||
tolerance = self.config.get("lane_tolerance", 3.0)
|
||||
|
||||
for idx, car in enumerate(self.car_birth_info_list):
|
||||
if self.debug_lane_filter:
|
||||
print(f"\n车辆 {idx+1}/{len(self.car_birth_info_list)}: ID={car['id']}")
|
||||
|
||||
if self._is_position_on_lane(car['begin'], tolerance=tolerance):
|
||||
valid_cars.append(car)
|
||||
if self.debug_lane_filter:
|
||||
print(f" ✅ 保留")
|
||||
else:
|
||||
filtered_cars.append({
|
||||
'id': car['id'],
|
||||
'position': car['begin'],
|
||||
'reason': '生成位置不在有效车道上(可能在草坪/停车场等区域)'
|
||||
if track["type"] == MetaDriveType.VEHICLE:
|
||||
_obj_to_clean_this_frame.append(scenario_id)
|
||||
valid = track["state"]["valid"]
|
||||
first_show = int(np.argmax(valid)) if valid.any() else -1
|
||||
last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else -1
|
||||
car_birth_info_list.append({
|
||||
"id": track["metadata"]["object_id"],
|
||||
"show_time": first_show,
|
||||
"begin": (track["state"]["position"][first_show, 0], track["state"]["position"][first_show, 1]),
|
||||
"heading": track["state"]["heading"][first_show],
|
||||
"end": (track["state"]["position"][last_show, 0], track["state"]["position"][last_show, 1]),
|
||||
})
|
||||
if self.debug_lane_filter:
|
||||
print(f" ❌ 过滤 (原因: 不在车道上)")
|
||||
|
||||
self.car_birth_info_list = valid_cars
|
||||
|
||||
if self.debug_lane_filter:
|
||||
print(f"\n📊 过滤结果: 保留 {len(valid_cars)} 辆, 过滤 {len(filtered_cars)} 辆")
|
||||
|
||||
return len(valid_cars), len(filtered_cars), filtered_cars
|
||||
|
||||
return car_birth_info_list, {}, _obj_to_clean_this_frame
|
||||
|
||||
def _spawn_controlled_agents(self):
|
||||
# ego_vehicle = self.engine.agent_manager.active_agents.get("default_agent")
|
||||
# ego_position = ego_vehicle.position if ego_vehicle else np.array([0, 0])
|
||||
@@ -299,148 +169,26 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
# ✅ 关键:注册到引擎的 active_agents,才能参与物理更新
|
||||
self.engine.agent_manager.active_agents[agent_id] = vehicle
|
||||
|
||||
def _get_traffic_light_state(self, vehicle):
|
||||
"""
|
||||
获取车辆当前位置的红绿灯状态(优化版)
|
||||
|
||||
解决问题:
|
||||
1. 部分红绿灯状态为None的问题 - 添加异常处理和默认值
|
||||
2. 车道分段导致无法获取红绿灯的问题 - 优先使用导航模块,失败时回退到遍历
|
||||
|
||||
Returns:
|
||||
int: 0=无红绿灯, 1=绿灯, 2=黄灯, 3=红灯
|
||||
"""
|
||||
traffic_light = 0
|
||||
state = vehicle.get_state()
|
||||
position_2d = state['position'][:2]
|
||||
|
||||
if self.debug_traffic_light:
|
||||
print(f"\n🚦 检测车辆红绿灯 - 位置: ({position_2d[0]:.1f}, {position_2d[1]:.1f})")
|
||||
|
||||
try:
|
||||
# 方法1:优先尝试从车辆导航模块获取当前车道(更高效)
|
||||
if hasattr(vehicle, 'navigation') and vehicle.navigation is not None:
|
||||
current_lane = vehicle.navigation.current_lane
|
||||
|
||||
if self.debug_traffic_light:
|
||||
print(f" 方法1-导航模块:")
|
||||
print(f" current_lane = {current_lane}")
|
||||
print(f" lane_index = {current_lane.index if current_lane else 'None'}")
|
||||
|
||||
if current_lane:
|
||||
has_light = self.engine.light_manager.has_traffic_light(current_lane.index)
|
||||
|
||||
if self.debug_traffic_light:
|
||||
print(f" has_traffic_light = {has_light}")
|
||||
|
||||
if has_light:
|
||||
status = self.engine.light_manager._lane_index_to_obj[current_lane.index].status
|
||||
|
||||
if self.debug_traffic_light:
|
||||
print(f" status = {status}")
|
||||
|
||||
if status == 'TRAFFIC_LIGHT_GREEN':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法1成功: 绿灯")
|
||||
return 1
|
||||
elif status == 'TRAFFIC_LIGHT_YELLOW':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法1成功: 黄灯")
|
||||
return 2
|
||||
elif status == 'TRAFFIC_LIGHT_RED':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法1成功: 红灯")
|
||||
return 3
|
||||
elif status is None:
|
||||
if self.debug_traffic_light:
|
||||
print(f" ⚠️ 方法1: 红绿灯状态为None")
|
||||
return 0
|
||||
else:
|
||||
if self.debug_traffic_light:
|
||||
print(f" 该车道没有红绿灯")
|
||||
else:
|
||||
if self.debug_traffic_light:
|
||||
print(f" 导航模块current_lane为None")
|
||||
else:
|
||||
if self.debug_traffic_light:
|
||||
has_nav = hasattr(vehicle, 'navigation')
|
||||
nav_not_none = vehicle.navigation is not None if has_nav else False
|
||||
print(f" 方法1-导航模块: 不可用 (hasattr={has_nav}, not_none={nav_not_none})")
|
||||
|
||||
except Exception as e:
|
||||
if self.debug_traffic_light:
|
||||
print(f" ❌ 方法1异常: {type(e).__name__}: {e}")
|
||||
pass
|
||||
|
||||
try:
|
||||
# 方法2:遍历所有车道查找(兜底方案,处理车道分段问题)
|
||||
if self.debug_traffic_light:
|
||||
print(f" 方法2-遍历车道: 开始遍历 {len(self.lanes)} 条车道")
|
||||
|
||||
found_lane = False
|
||||
checked_lanes = 0
|
||||
|
||||
for lane in self.lanes.values():
|
||||
try:
|
||||
checked_lanes += 1
|
||||
if lane.lane.point_on_lane(position_2d):
|
||||
found_lane = True
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✓ 找到车辆所在车道: {lane.lane.index} (检查了{checked_lanes}条)")
|
||||
|
||||
has_light = self.engine.light_manager.has_traffic_light(lane.lane.index)
|
||||
if self.debug_traffic_light:
|
||||
print(f" has_traffic_light = {has_light}")
|
||||
|
||||
if has_light:
|
||||
status = self.engine.light_manager._lane_index_to_obj[lane.lane.index].status
|
||||
if self.debug_traffic_light:
|
||||
print(f" status = {status}")
|
||||
|
||||
if status == 'TRAFFIC_LIGHT_GREEN':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法2成功: 绿灯")
|
||||
return 1
|
||||
elif status == 'TRAFFIC_LIGHT_YELLOW':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法2成功: 黄灯")
|
||||
return 2
|
||||
elif status == 'TRAFFIC_LIGHT_RED':
|
||||
if self.debug_traffic_light:
|
||||
print(f" ✅ 方法2成功: 红灯")
|
||||
return 3
|
||||
elif status is None:
|
||||
if self.debug_traffic_light:
|
||||
print(f" ⚠️ 方法2: 红绿灯状态为None")
|
||||
return 0
|
||||
else:
|
||||
if self.debug_traffic_light:
|
||||
print(f" 该车道没有红绿灯")
|
||||
break
|
||||
except:
|
||||
continue
|
||||
|
||||
if self.debug_traffic_light and not found_lane:
|
||||
print(f" ⚠️ 未找到车辆所在车道 (检查了{checked_lanes}条)")
|
||||
|
||||
except Exception as e:
|
||||
if self.debug_traffic_light:
|
||||
print(f" ❌ 方法2异常: {type(e).__name__}: {e}")
|
||||
pass
|
||||
|
||||
if self.debug_traffic_light:
|
||||
print(f" 结果: 返回 {traffic_light} (无红绿灯/未知)")
|
||||
|
||||
return traffic_light
|
||||
|
||||
def _get_all_obs(self):
|
||||
# position, velocity, heading, lidar, navigation, TODO: trafficlight -> list
|
||||
self.obs_list = []
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
state = vehicle.get_state()
|
||||
|
||||
# 使用优化后的红绿灯检测方法
|
||||
traffic_light = self._get_traffic_light_state(vehicle)
|
||||
traffic_light = 0
|
||||
for lane in self.lanes.values():
|
||||
if lane.lane.point_on_lane(state['position'][:2]):
|
||||
if self.engine.light_manager.has_traffic_light(lane.lane.index):
|
||||
traffic_light = self.engine.light_manager._lane_index_to_obj[lane.lane.index].status
|
||||
if traffic_light == 'TRAFFIC_LIGHT_GREEN':
|
||||
traffic_light = 1
|
||||
elif traffic_light == 'TRAFFIC_LIGHT_YELLOW':
|
||||
traffic_light = 2
|
||||
elif traffic_light == 'TRAFFIC_LIGHT_RED':
|
||||
traffic_light = 3
|
||||
else:
|
||||
traffic_light = 0
|
||||
break
|
||||
|
||||
lidar = self.engine.get_sensor("lidar").perceive(num_lasers=80, distance=30, base_vehicle=vehicle,
|
||||
physics_world=self.engine.physics_world.dynamic_world)
|
||||
@@ -465,6 +213,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:
|
||||
@@ -476,4 +225,4 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
||||
dones = {aid: False for aid in self.controlled_agents}
|
||||
dones["__all__"] = self.episode_step >= self.config["horizon"]
|
||||
infos = {aid: {} for aid in self.controlled_agents}
|
||||
return obs, rewards, dones, infos
|
||||
return obs, rewards, dones, infos
|
||||
@@ -1,219 +0,0 @@
|
||||
"""
|
||||
测试车道过滤和红绿灯检测功能
|
||||
"""
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from simple_idm_policy import ConstantVelocityPolicy
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
from logger_utils import setup_logger
|
||||
import os
|
||||
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/Env"
|
||||
|
||||
def test_lane_filter():
|
||||
"""测试车道过滤功能(基础版)"""
|
||||
print("=" * 60)
|
||||
print("测试1:车道过滤功能(基础)")
|
||||
print("=" * 60)
|
||||
|
||||
# 创建启用过滤的环境
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 100,
|
||||
"use_render": False,
|
||||
|
||||
# 车道过滤配置
|
||||
"filter_offroad_vehicles": True,
|
||||
"lane_tolerance": 3.0,
|
||||
"max_controlled_vehicles": 10,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
print("\n启用车道过滤...")
|
||||
obs = env.reset(0)
|
||||
print(f"生成车辆数: {len(env.controlled_agents)}")
|
||||
print(f"观测数据长度: {len(obs)}")
|
||||
|
||||
# 运行几步
|
||||
for step in range(5):
|
||||
actions = {aid: env.controlled_agents[aid].policy.act()
|
||||
for aid in env.controlled_agents}
|
||||
obs, rewards, dones, infos = env.step(actions)
|
||||
|
||||
env.close()
|
||||
print("✓ 车道过滤测试通过\n")
|
||||
|
||||
|
||||
def test_lane_filter_debug():
|
||||
"""测试车道过滤功能(详细调试)"""
|
||||
print("=" * 60)
|
||||
print("测试1b:车道过滤功能(详细调试模式)")
|
||||
print("=" * 60)
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 100,
|
||||
"use_render": False,
|
||||
|
||||
# 车道过滤配置
|
||||
"filter_offroad_vehicles": True,
|
||||
"lane_tolerance": 3.0,
|
||||
"max_controlled_vehicles": 5, # 只看前5辆车
|
||||
|
||||
# 🔥 启用调试模式
|
||||
"debug_lane_filter": True, # 启用车道过滤调试
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
print("\n启用车道过滤调试...")
|
||||
obs = env.reset(0)
|
||||
|
||||
env.close()
|
||||
print("\n✓ 车道过滤调试测试完成\n")
|
||||
|
||||
|
||||
def test_traffic_light():
|
||||
"""测试红绿灯检测功能"""
|
||||
print("=" * 60)
|
||||
print("测试2:红绿灯检测功能(启用详细调试)")
|
||||
print("=" * 60)
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 100,
|
||||
"use_render": False,
|
||||
"filter_offroad_vehicles": True,
|
||||
"max_controlled_vehicles": 3, # 只测试3辆车
|
||||
|
||||
# 🔥 启用调试模式
|
||||
"debug_traffic_light": True, # 启用红绿灯调试
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
obs = env.reset(0)
|
||||
|
||||
# 测试红绿灯检测(调试模式会自动输出详细信息)
|
||||
print(f"\n" + "="*60)
|
||||
print(f"开始逐车检测红绿灯状态(共 {len(env.controlled_agents)} 辆车)")
|
||||
print("="*60)
|
||||
|
||||
for idx, (aid, vehicle) in enumerate(list(env.controlled_agents.items())[:3]): # 只测试前3辆
|
||||
print(f"\n【车辆 {idx+1}/3】 ID={aid}")
|
||||
traffic_light = env._get_traffic_light_state(vehicle)
|
||||
state = vehicle.get_state()
|
||||
|
||||
status_text = {0: '无/未知', 1: '绿灯', 2: '黄灯', 3: '红灯'}[traffic_light]
|
||||
print(f"最终结果: 红绿灯状态={traffic_light} ({status_text})\n")
|
||||
|
||||
env.close()
|
||||
print("="*60)
|
||||
print("✓ 红绿灯检测测试完成")
|
||||
print("="*60 + "\n")
|
||||
|
||||
|
||||
def test_without_filter():
|
||||
"""测试禁用过滤的情况"""
|
||||
print("=" * 60)
|
||||
print("测试3:禁用过滤(对比测试)")
|
||||
print("=" * 60)
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": AssetLoader.file_path(WAYMO_DATA_DIR, "exp_converted", unix_style=False),
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": 100,
|
||||
"use_render": False,
|
||||
|
||||
# 禁用过滤
|
||||
"filter_offroad_vehicles": False,
|
||||
"max_controlled_vehicles": None,
|
||||
},
|
||||
agent2policy=ConstantVelocityPolicy(target_speed=50)
|
||||
)
|
||||
|
||||
print("\n禁用车道过滤...")
|
||||
obs = env.reset(0)
|
||||
print(f"生成车辆数(未过滤): {len(env.controlled_agents)}")
|
||||
|
||||
env.close()
|
||||
print("✓ 禁用过滤测试通过\n")
|
||||
|
||||
|
||||
def run_tests(debug_mode=False):
|
||||
"""运行测试的主函数"""
|
||||
try:
|
||||
if debug_mode:
|
||||
print("🐛 调试模式启用")
|
||||
print("=" * 60 + "\n")
|
||||
test_lane_filter_debug()
|
||||
test_traffic_light()
|
||||
else:
|
||||
print("⚡ 标准测试模式(使用 --debug 参数启用详细调试)")
|
||||
print("=" * 60 + "\n")
|
||||
test_lane_filter()
|
||||
test_traffic_light()
|
||||
test_without_filter()
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("✅ 所有测试通过!")
|
||||
print("=" * 60)
|
||||
print("\n功能说明:")
|
||||
print("1. 车道过滤功能已启用,自动过滤非车道区域车辆")
|
||||
print("2. 红绿灯检测采用双重策略,确保稳定获取状态")
|
||||
print("3. 可通过配置参数灵活启用/禁用功能")
|
||||
print("\n使用方法:")
|
||||
print(" python Env/test_lane_filter.py # 标准测试")
|
||||
print(" python Env/test_lane_filter.py --debug # 详细调试")
|
||||
print(" python Env/test_lane_filter.py --log # 保存日志")
|
||||
print(" python Env/test_lane_filter.py --debug --log # 调试+日志")
|
||||
print("\n请运行 run_multiagent_env.py 查看完整效果")
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n❌ 测试失败: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
# 解析命令行参数
|
||||
debug_mode = "--debug" in sys.argv or "-d" in sys.argv
|
||||
enable_logging = "--log" in sys.argv or "-l" in sys.argv
|
||||
|
||||
# 提取自定义日志文件名
|
||||
log_file = None
|
||||
for arg in sys.argv:
|
||||
if arg.startswith("--log-file="):
|
||||
log_file = arg.split("=")[1]
|
||||
break
|
||||
|
||||
if enable_logging:
|
||||
# 启用日志记录
|
||||
log_dir = os.path.join(os.path.dirname(__file__), "logs")
|
||||
|
||||
# 生成默认日志文件名
|
||||
if log_file is None:
|
||||
mode_suffix = "debug" if debug_mode else "standard"
|
||||
from datetime import datetime
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
log_file = f"test_{mode_suffix}_{timestamp}.log"
|
||||
|
||||
with setup_logger(log_file=log_file, log_dir=log_dir):
|
||||
run_tests(debug_mode=debug_mode)
|
||||
else:
|
||||
# 不启用日志,直接运行
|
||||
run_tests(debug_mode=debug_mode)
|
||||
|
||||
221
Env/utils.py
221
Env/utils.py
@@ -2,6 +2,227 @@ import numpy as np
|
||||
import torch
|
||||
import random
|
||||
|
||||
from metadrive.type import MetaDriveType
|
||||
|
||||
|
||||
def _static_obbs_overlap(car_a, car_b):
|
||||
"""
|
||||
Check if two static vehicle OBBs overlap (2D SAT).
|
||||
car_a, car_b: dicts with "begin" (x, y), "heading" (rad), "length", "width".
|
||||
begin is center; half-extents are length/2, width/2.
|
||||
"""
|
||||
def _get_corners(car):
|
||||
cx, cy = car["begin"][0], car["begin"][1]
|
||||
h = float(car["heading"])
|
||||
L2 = float(car["length"]) / 2.0
|
||||
W2 = float(car["width"]) / 2.0
|
||||
ux, uy = np.cos(h), np.sin(h)
|
||||
vx, vy = -np.sin(h), np.cos(h)
|
||||
return np.array([
|
||||
[cx + L2 * ux + W2 * vx, cy + L2 * uy + W2 * vy],
|
||||
[cx + L2 * ux - W2 * vx, cy + L2 * uy - W2 * vy],
|
||||
[cx - L2 * ux - W2 * vx, cy - L2 * uy - W2 * vy],
|
||||
[cx - L2 * ux + W2 * vx, cy - L2 * uy + W2 * vy],
|
||||
])
|
||||
|
||||
def _get_axes(car):
|
||||
h = float(car["heading"])
|
||||
return [
|
||||
np.array([np.cos(h), np.sin(h)]),
|
||||
np.array([-np.sin(h), np.cos(h)]),
|
||||
]
|
||||
|
||||
corners_a = _get_corners(car_a)
|
||||
corners_b = _get_corners(car_b)
|
||||
axes = _get_axes(car_a) + _get_axes(car_b)
|
||||
|
||||
for axis in axes:
|
||||
proj_a = corners_a @ axis
|
||||
proj_b = corners_b @ axis
|
||||
min_a, max_a = proj_a.min(), proj_a.max()
|
||||
min_b, max_b = proj_b.min(), proj_b.max()
|
||||
if max_a < min_b or max_b < min_a:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _deduplicate_background_by_collision(background_vehicles):
|
||||
"""
|
||||
Merge static tracks that collide (same physical vehicle). Build collision graph,
|
||||
find connected components, keep one representative per component (min show_time, then scenario_id).
|
||||
"""
|
||||
if not background_vehicles:
|
||||
return background_vehicles
|
||||
items = list(background_vehicles.items())
|
||||
n = len(items)
|
||||
# Build adjacency by index
|
||||
parent = list(range(n))
|
||||
|
||||
def find(i):
|
||||
if parent[i] != i:
|
||||
parent[i] = find(parent[i])
|
||||
return parent[i]
|
||||
|
||||
def union(i, j):
|
||||
pi, pj = find(i), find(j)
|
||||
if pi != pj:
|
||||
parent[pi] = pj
|
||||
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
if _static_obbs_overlap(items[i][1], items[j][1]):
|
||||
union(i, j)
|
||||
|
||||
# Representative per component: index with min (show_time, scenario_id)
|
||||
comp_rep = {}
|
||||
for i in range(n):
|
||||
r = find(i)
|
||||
sid, car = items[i][0], items[i][1]
|
||||
key = (car.get("show_time", 0), sid)
|
||||
if r not in comp_rep or key < comp_rep[r][0]:
|
||||
comp_rep[r] = (key, sid, car)
|
||||
|
||||
return {sid: car for (_, sid, car) in comp_rep.values()}
|
||||
|
||||
|
||||
def is_on_lane(pos, map_manager, threshold=2.0):
|
||||
"""Check if a position is on a valid lane (within lateral tolerance)."""
|
||||
if map_manager is None or map_manager.current_map is None:
|
||||
return True
|
||||
try:
|
||||
lane, _ = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True)
|
||||
if lane is None:
|
||||
return False
|
||||
long, lat = lane.local_coordinates(pos)
|
||||
width = lane.width
|
||||
if abs(lat) <= (width / 2 + threshold):
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def filter_traffic_tracks_to_birth_lists(
|
||||
current_traffic_data,
|
||||
sdc_scenario_id,
|
||||
map_manager,
|
||||
*,
|
||||
lane_threshold=5.0,
|
||||
static_displacement_threshold=5.0,
|
||||
static_speed_threshold=1.0,
|
||||
return_stats=False,
|
||||
deduplicate_static_by_collision=True,
|
||||
):
|
||||
"""
|
||||
Filter traffic tracks into controlled (car_birth_info_list) and background lists.
|
||||
|
||||
- controlled (car_birth_info_list): 非 SDC、类型 VEHICLE、至少一帧 valid、在车道内、且非静态
|
||||
(位移/速度超过阈值)。用于策略控制或专家回放,spawn 时机为 show_time == round。
|
||||
- background (background_vehicles): 同上但在车道内且判定为静态(位移 < 5m、速度 < 1 m/s)。
|
||||
仅作场景占位与观测邻居,spawn 时机为 show_time == round,按 valid 在 step 中移除。
|
||||
|
||||
Returns (car_birth_info_list, background_vehicles, obj_to_clean) or, if return_stats=True,
|
||||
(car_birth_info_list, background_vehicles, obj_to_clean, stats_dict).
|
||||
stats_dict: n_total, n_no_valid, n_off_lane, n_static, n_controlled.
|
||||
"""
|
||||
car_birth_info_list = []
|
||||
background_vehicles = {}
|
||||
obj_to_clean = []
|
||||
n_total = 0
|
||||
n_no_valid = 0
|
||||
n_off_lane = 0
|
||||
n_static = 0
|
||||
|
||||
for scenario_id, track in current_traffic_data.items():
|
||||
if scenario_id == sdc_scenario_id:
|
||||
continue
|
||||
if track["type"] != MetaDriveType.VEHICLE:
|
||||
continue
|
||||
|
||||
n_total += 1
|
||||
obj_to_clean.append(scenario_id)
|
||||
valid = track["state"]["valid"]
|
||||
if not valid.any():
|
||||
n_no_valid += 1
|
||||
continue
|
||||
|
||||
first_show = int(np.argmax(valid))
|
||||
last_show = len(valid) - 1 - int(np.argmax(valid[::-1]))
|
||||
mid_show = (first_show + last_show) // 2
|
||||
|
||||
start_pos = track["state"]["position"][first_show]
|
||||
is_valid_track = True
|
||||
if not is_on_lane(start_pos, map_manager, threshold=lane_threshold):
|
||||
mid_pos = track["state"]["position"][mid_show]
|
||||
if not is_on_lane(mid_pos, map_manager, threshold=lane_threshold):
|
||||
is_valid_track = False
|
||||
|
||||
if not is_valid_track:
|
||||
n_off_lane += 1
|
||||
continue
|
||||
|
||||
positions = track["state"]["position"][valid.astype(bool)]
|
||||
velocities = track["state"]["velocity"][valid.astype(bool)]
|
||||
total_displacement = 0.0
|
||||
max_speed = 0.0
|
||||
if len(positions) > 1:
|
||||
total_displacement = float(np.linalg.norm(positions[-1] - positions[0]))
|
||||
max_speed = float(np.max(np.linalg.norm(velocities, axis=1)))
|
||||
is_static = total_displacement < static_displacement_threshold and max_speed < static_speed_threshold
|
||||
|
||||
if is_static:
|
||||
n_static += 1
|
||||
background_vehicles[scenario_id] = {
|
||||
"id": track["metadata"]["object_id"],
|
||||
"show_time": first_show,
|
||||
"begin": (
|
||||
float(track["state"]["position"][first_show, 0]),
|
||||
float(track["state"]["position"][first_show, 1]),
|
||||
),
|
||||
"heading": float(track["state"]["heading"][first_show]),
|
||||
"end": (
|
||||
float(track["state"]["position"][last_show, 0]),
|
||||
float(track["state"]["position"][last_show, 1]),
|
||||
),
|
||||
"scenario_id": scenario_id,
|
||||
"length": track["state"]["length"][first_show],
|
||||
"width": track["state"]["width"][first_show],
|
||||
"valid": valid,
|
||||
}
|
||||
continue
|
||||
|
||||
car_birth_info_list.append({
|
||||
"id": track["metadata"]["object_id"],
|
||||
"show_time": first_show,
|
||||
"begin": (
|
||||
float(track["state"]["position"][first_show, 0]),
|
||||
float(track["state"]["position"][first_show, 1]),
|
||||
),
|
||||
"heading": float(track["state"]["heading"][first_show]),
|
||||
"end": (
|
||||
float(track["state"]["position"][last_show, 0]),
|
||||
float(track["state"]["position"][last_show, 1]),
|
||||
),
|
||||
"scenario_id": scenario_id,
|
||||
"length": track["state"]["length"][first_show],
|
||||
"width": track["state"]["width"][first_show],
|
||||
})
|
||||
|
||||
if deduplicate_static_by_collision and background_vehicles:
|
||||
background_vehicles = _deduplicate_background_by_collision(background_vehicles)
|
||||
|
||||
if return_stats:
|
||||
stats = {
|
||||
"n_total": n_total,
|
||||
"n_no_valid": n_no_valid,
|
||||
"n_off_lane": n_off_lane,
|
||||
"n_static": n_static,
|
||||
"n_controlled": len(car_birth_info_list),
|
||||
}
|
||||
return car_birth_info_list, background_vehicles, obj_to_clean, stats
|
||||
return car_birth_info_list, background_vehicles, obj_to_clean
|
||||
|
||||
|
||||
def set_seed(seed):
|
||||
if seed == -1:
|
||||
seed = np.random.randint(0, 10000)
|
||||
|
||||
201
README.md
201
README.md
@@ -1,85 +1,148 @@
|
||||
# MAGAIL4AutoDrive
|
||||
### 1.1 环境搭建
|
||||
环境核心代码封装于`Env`文件夹,通过运行`run_multiagent_env.py`即可启动多智能体交互环境,该脚本的核心功能为读取各智能体(车辆)的动作指令,并将其传入`env.step()`方法中完成仿真执行。
|
||||
|
||||
**性能优化版本:** 针对原始版本FPS低(15帧)和CPU利用率不足的问题,已提供多个优化版本:
|
||||
- `run_multiagent_env_fast.py` - 激光雷达优化版(30-60 FPS,2-4倍提升)⭐推荐
|
||||
- `run_multiagent_env_parallel.py` - 多进程并行版(300-600 steps/s总吞吐量,充分利用多核CPU)⭐⭐推荐
|
||||
- 详见 `Env/QUICK_START.md` 快速使用指南
|
||||
基于 **MetaDrive** 仿真器和 **Waymo Open Motion Dataset** 的自动驾驶多智能体模仿学习(MAGAIL)与行为克隆(BC)训练系统。
|
||||
|
||||
当前已初步实现`Env.senario_env.MultiAgentScenarioEnv.reset()`车辆生成函数,具体逻辑如下:首先读取专家数据集中各车辆的初始位姿信息;随后对原始数据进行清洗,剔除车辆 Agent 实例信息,记录核心参数(车辆 ID、初始生成位置、朝向角、生成时间戳、目标终点坐标);最后调用`_spawn_controlled_agents()`函数,依据清洗后的参数在指定时间、指定位置生成搭载自动驾驶算法的可控车辆。
|
||||
本项目旨在从真实的 Waymo 驾驶数据中提取专家轨迹,并通过模仿学习(Imitation Learning)训练能够适应复杂交互场景的自动驾驶策略。
|
||||
|
||||
**✅ 已解决:车辆生成位置偏差问题**
|
||||
- **问题描述**:部分车辆生成于草坪、停车场等非车道区域,原因是专家数据记录误差或停车场特殊标注
|
||||
- **解决方案**:实现了`_is_position_on_lane()`车道区域检测机制和`_filter_valid_spawn_positions()`过滤函数
|
||||
- 检测逻辑:通过`point_on_lane()`判断位置是否在车道上,支持容差参数(默认3米)处理边界情况
|
||||
- 双重检测:优先使用精确检测,失败时使用容差范围检测,确保车道边缘车辆不被误过滤
|
||||
- 自动过滤:在`reset()`时自动过滤非车道区域车辆,并输出过滤统计信息
|
||||
- **配置参数**:
|
||||
- `filter_offroad_vehicles=True`:启用/禁用车道过滤功能
|
||||
- `lane_tolerance=3.0`:车道检测容差(米),可根据场景调整
|
||||
- `max_controlled_vehicles=10`:限制最大车辆数(可选)
|
||||
- **使用示例**:在环境配置中设置上述参数即可自动启用,运行时会显示过滤信息(如"过滤5辆,保留45辆")
|
||||
## 目录结构
|
||||
|
||||
```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 共用)
|
||||
│ ├── bc_ego_replay_env.py # BCEgoReplayEnv,单智能体 BC 评估(仅 ego 受控)
|
||||
│ ├── 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
|
||||
```
|
||||
|
||||
### 1.2 观测获取
|
||||
观测信息采集功能通过`Env.senario_env.MultiAgentScenarioEnv._get_all_obs()`函数实现,该函数支持遍历所有可控车辆并采集多维度观测数据,当前已实现的观测维度包括:车辆实时位置坐标、朝向角、行驶速度、雷达扫描点云(含障碍物与车道线特征)、导航信息(因场景复杂度较低,暂采用目标终点坐标直接作为导航输入)。
|
||||
## 路径约定(相对项目根)
|
||||
|
||||
**✅ 已解决:红绿灯信息采集问题**
|
||||
- **问题描述**:
|
||||
- 问题1:部分红绿灯状态值为`None`,导致异常或错误判断
|
||||
- 问题2:车道分段设计时,部分区域车辆无法匹配到红绿灯
|
||||
- **解决方案**:实现了`_get_traffic_light_state()`优化方法,采用多级检测策略
|
||||
- **方法1(优先)**:从车辆导航模块`vehicle.navigation.current_lane`获取当前车道,直接查询红绿灯状态(高效,自动处理车道分段)
|
||||
- **方法2(兜底)**:遍历所有车道,通过`point_on_lane()`判断车辆位置,查找对应红绿灯(处理导航失败情况)
|
||||
- **异常处理**:对状态为`None`的情况返回0(无红绿灯),所有异常均有try-except保护,确保不会中断程序
|
||||
- **返回值规范**:0=无红绿灯/未知, 1=绿灯, 2=黄灯, 3=红灯
|
||||
- **优势**:双重保障机制,优先用高效方法,失败时自动切换到兜底方案,确保所有场景都能正确获取红绿灯信息
|
||||
- **数据**:Waymo 场景 `data/exp_filtered`;专家 pkl `data/training_data`;其他轨迹 `data/trajectories`
|
||||
- **模型**:BC `models/bc/`,MAGAIL `models/magail/`
|
||||
- **日志**:TensorBoard 写入 `logs/bc/`、`logs/magail/`
|
||||
|
||||
所有默认路径均为相对项目根,便于在不同设备上复用。
|
||||
|
||||
### 1.3 算法模块
|
||||
本方案的核心创新点在于对 GAIL 算法的判别器进行改进,使其适配多智能体场景下 “输入长度动态变化”(车辆数量不固定)的特性,实现对整体交互场景的分类判断,进而满足多智能体自动驾驶环境的训练需求。算法核心代码封装于`Algorithm.bert.Bert`类,具体实现逻辑如下:
|
||||
## 数据处理流程
|
||||
|
||||
1. 输入层处理:输入数据为维度`(N, input_dim)`的矩阵(其中`N`为当前场景车辆数量,`input_dim`为单车辆固定观测维度),初始化`Bert`类时需设置`input_dim`,确保输入维度匹配;
|
||||
2. 嵌入层与位置编码:通过`projection`线性投影层将单车辆观测维度映射至预设的嵌入维度(`embed_dim`),随后叠加可学习的位置编码(`pos_embed`),以捕捉观测序列的时序与空间关联信息;
|
||||
3. Transformer 特征提取:嵌入后的特征向量输入至多层`Transformer`网络(层数由`num_layers`参数控制),完成高阶特征交互与抽象;
|
||||
4. 分类头设计:提供两种特征聚合与分类方案:若开启`CLS`模式,在嵌入层前拼接 1 个可学习的`CLS`标记,最终取`CLS`标记对应的特征向量输入全连接层完成分类;若关闭`CLS`模式,则对`Transformer`输出的所有车辆特征向量进行序列维度均值池化,再将池化后的全局特征输入全连接层。分类器支持可选的`Tanh`激活函数,以适配不同场景下的输出分布需求。
|
||||
从 Waymo Motion 原始数据到本项目训练用专家 pkl,依次为:
|
||||
|
||||
**1) 下载 Waymo Motion(TFRecord)**
|
||||
安装 `gsutil` 并登录 Google 账号后,例如只下载 training_20s:
|
||||
|
||||
### 1.4 动作执行
|
||||
在当前环境测试阶段,暂沿用腾达的动作执行框架:为每辆可控车辆分配独立的`policy`模型,将单车辆观测数据输入对应`policy`得到动作指令后,传入`env.step()`完成仿真;同时在`before_step`阶段调用`_set_action()`函数,将动作指令绑定至车辆实例,最终由 MetaDrive 仿真系统完成物理动力学计算与场景渲染。
|
||||
|
||||
后续优化方向为构建 "参数共享式统一模型框架",具体设计如下:所有车辆共用 1 个`policy`模型,通过参数共享机制实现模型的全局统一维护。该框架具备三重优势:一是避免多车辆独立模型带来的训练偏差(如不同模型训练程度不一致);二是解决车辆数量动态变化时的模型管理问题(车辆新增无需额外初始化模型,车辆减少不丢失模型训练信息);三是支持动作指令的并行计算,可显著提升每一步决策的迭代效率,适配大规模多智能体交互场景的训练需求。
|
||||
|
||||
---
|
||||
|
||||
## 问题解决总结
|
||||
|
||||
### ✅ 已完成的优化
|
||||
|
||||
1. **车辆生成位置偏差** - 实现车道区域检测和自动过滤,配置参数:`filter_offroad_vehicles`, `lane_tolerance`, `max_controlled_vehicles`
|
||||
2. **红绿灯信息采集** - 采用双重检测策略(导航模块+遍历兜底),处理None状态和车道分段问题
|
||||
3. **性能优化** - 提供多个优化版本(fast/parallel),FPS从15提升到30-60,支持多进程充分利用CPU
|
||||
|
||||
### 🧪 测试方法
|
||||
```bash
|
||||
# 测试车道过滤和红绿灯检测
|
||||
python Env/test_lane_filter.py
|
||||
|
||||
# 运行标准版本(带过滤)
|
||||
python Env/run_multiagent_env.py
|
||||
|
||||
# 运行高性能版本
|
||||
python Env/run_multiagent_env_fast.py
|
||||
gsutil -m cp -r "gs://waymo_open_dataset_motion_v_1_2_0/uncompressed/scenario/training_20s" ./waymo/
|
||||
```
|
||||
|
||||
### 📝 配置示例
|
||||
```python
|
||||
config = {
|
||||
# 车道过滤
|
||||
"filter_offroad_vehicles": True, # 启用车道过滤
|
||||
"lane_tolerance": 3.0, # 容差范围(米)
|
||||
"max_controlled_vehicles": 10, # 最大车辆数
|
||||
# 其他配置...
|
||||
}
|
||||
**2) ScenarioNet Convert(TFRecord → ScenarioNet 场景库)**
|
||||
需安装 ScenarioNet、MetaDrive 及 TensorFlow 2.11、protobuf 3.20;转换时不用 GPU。
|
||||
|
||||
```bash
|
||||
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)。
|
||||
|
||||
**4) 本项目:生成专家 pkl**
|
||||
使用筛选后的场景目录,生成训练用 pkl 到 `data/training_data`:
|
||||
|
||||
- **多智能体**(所有受控车轨迹,输出 `expert_data_{start_index}_{num_scenarios}.pkl`):
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||
```
|
||||
|
||||
- **单智能体**(仅 ego 车轨迹,输出 `expert_data_ego_{start_index}_{num_scenarios}.pkl`,用于单智能体 BC):
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||
```
|
||||
|
||||
## 核心工作流
|
||||
|
||||
### 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 --start_index 0
|
||||
```
|
||||
|
||||
- **单智能体(仅 ego)**:
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||
```
|
||||
|
||||
### 2. 行为克隆 (BC)
|
||||
BC 支持两种模式:**多智能体**(默认,所有受控车共用同一策略)与 **单智能体**(仅 ego 车,评估时其他车按专家轨迹回放)。
|
||||
|
||||
- **多智能体训练**(模型保存到 `models/bc/`,日志到 `logs/bc/`):
|
||||
```bash
|
||||
python train_bc.py --expert_data_path data/training_data/expert_data_0_50.pkl --epochs 100
|
||||
```
|
||||
|
||||
- **单智能体训练**(使用 ego-only 数据,评估时仅 ego 受策略控制,其他车专家回放):
|
||||
```bash
|
||||
python train_bc.py --expert_data_path data/training_data/expert_data_ego_0_50.pkl --epochs 100 --single_agent
|
||||
```
|
||||
|
||||
- **可视化**:`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt`
|
||||
仅自车用策略、其他车回放(单智能体可视化):加 `--ego_only`,例如
|
||||
`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --ego_only --num_scenarios 1`
|
||||
|
||||
### 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/bc_ego_replay_env.py**:`BCEgoReplayEnv`,单智能体 BC 评估环境,仅 ego 受策略控制,其他车按专家轨迹回放
|
||||
- **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)。
|
||||
|
||||
2
algorithms/__init__.py
Normal file
2
algorithms/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
||||
"""Compatibility package for legacy HBBC checkpoints."""
|
||||
|
||||
18
algorithms/utils.py
Normal file
18
algorithms/utils.py
Normal file
@@ -0,0 +1,18 @@
|
||||
import numpy as np
|
||||
|
||||
|
||||
class RunningMeanStd(object):
|
||||
def __init__(self, epsilon=1e-4, shape=()):
|
||||
self.mean = np.zeros(shape, np.float64)
|
||||
self.var = np.ones(shape, np.float64)
|
||||
self.count = epsilon
|
||||
|
||||
|
||||
class Normalizer(RunningMeanStd):
|
||||
def __init__(self, input_dim, epsilon=1e-4, clip_obs=10.0):
|
||||
super().__init__(shape=input_dim)
|
||||
self.epsilon = epsilon
|
||||
self.clip_obs = clip_obs
|
||||
|
||||
def normalize(self, input):
|
||||
return np.clip((input - self.mean) / np.sqrt(self.var + self.epsilon), -self.clip_obs, self.clip_obs)
|
||||
BIN
analysis_results/distributions.png
Normal file
BIN
analysis_results/distributions.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 316 KiB |
BIN
analysis_results/statistics.pkl
Normal file
BIN
analysis_results/statistics.pkl
Normal file
Binary file not shown.
0
dataset/__init__.py
Normal file
0
dataset/__init__.py
Normal file
BIN
dataset/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
dataset/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
dataset/__pycache__/__init__.cpython-39.pyc
Normal file
BIN
dataset/__pycache__/__init__.cpython-39.pyc
Normal file
Binary file not shown.
BIN
dataset/__pycache__/magail_dataset.cpython-313.pyc
Normal file
BIN
dataset/__pycache__/magail_dataset.cpython-313.pyc
Normal file
Binary file not shown.
BIN
dataset/__pycache__/magail_dataset.cpython-39.pyc
Normal file
BIN
dataset/__pycache__/magail_dataset.cpython-39.pyc
Normal file
Binary file not shown.
305
dataset/expert_dataset.py
Normal file
305
dataset/expert_dataset.py
Normal file
@@ -0,0 +1,305 @@
|
||||
import sys
|
||||
import os
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
project_root = os.path.dirname(current_dir)
|
||||
sys.path.insert(0, os.path.join(project_root, "Env"))
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
import pickle
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
|
||||
class DummyPolicy:
|
||||
def act(self, *args, **kwargs):
|
||||
return np.array([0.0, 0.0])
|
||||
|
||||
class ExpertTrajectoryDataset(Dataset):
|
||||
"""
|
||||
完整107维观测的专家轨迹数据集
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
trajectory_data: dict,
|
||||
observation_data: dict = None, # 可选的完整观测
|
||||
sequence_length: int = 1,
|
||||
extract_actions: bool = True):
|
||||
"""
|
||||
Args:
|
||||
trajectory_data: 专家轨迹数据
|
||||
observation_data: 完整107维观测数据(可选)
|
||||
sequence_length: 序列长度
|
||||
extract_actions: 是否提取动作
|
||||
"""
|
||||
self.trajectory_data = trajectory_data
|
||||
self.observation_data = observation_data if observation_data else {}
|
||||
self.sequence_length = sequence_length
|
||||
self.extract_actions = extract_actions
|
||||
|
||||
# 构建索引
|
||||
self.indices = []
|
||||
for traj_id, traj in trajectory_data.items():
|
||||
traj_len = traj["length"]
|
||||
for start_idx in range(traj_len - sequence_length):
|
||||
self.indices.append((traj_id, start_idx))
|
||||
|
||||
obs_dim = 107 if len(self.observation_data) > 0 else 5
|
||||
print(f"专家数据集: {len(trajectory_data)} 条轨迹, "
|
||||
f"{len(self.indices)} 个训练样本, 观测维度: {obs_dim}")
|
||||
|
||||
def __len__(self):
|
||||
return len(self.indices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
traj_id, start_idx = self.indices[idx]
|
||||
traj = self.trajectory_data[traj_id]
|
||||
|
||||
end_idx = start_idx + self.sequence_length
|
||||
|
||||
# 如果有完整观测,使用完整观测(107维)
|
||||
if traj_id in self.observation_data and len(self.observation_data[traj_id]) > 0:
|
||||
obs_sequence = self.observation_data[traj_id]
|
||||
states = obs_sequence[start_idx:end_idx] # (seq_len, 107)
|
||||
else:
|
||||
# 否则使用简化观测(5维)
|
||||
positions = traj["positions"][start_idx:end_idx+1]
|
||||
headings = traj["headings"][start_idx:end_idx+1]
|
||||
velocities = traj["velocities"][start_idx:end_idx]
|
||||
|
||||
states = []
|
||||
for i in range(self.sequence_length):
|
||||
state = np.concatenate([
|
||||
positions[i, :2], # x, y
|
||||
velocities[i], # vx, vy
|
||||
[headings[i]], # heading
|
||||
])
|
||||
states.append(state)
|
||||
states = np.array(states)
|
||||
|
||||
if self.extract_actions:
|
||||
positions = traj["positions"][start_idx:end_idx+1]
|
||||
headings = traj["headings"][start_idx:end_idx+1]
|
||||
velocities = traj["velocities"][start_idx:end_idx]
|
||||
|
||||
actions = self._extract_actions_from_states(
|
||||
positions[:-1], positions[1:],
|
||||
headings[:-1], headings[1:],
|
||||
velocities
|
||||
)
|
||||
return torch.FloatTensor(states), torch.FloatTensor(actions)
|
||||
else:
|
||||
next_states = states[1:]
|
||||
return torch.FloatTensor(states[:-1]), torch.FloatTensor(next_states)
|
||||
|
||||
def _extract_actions_from_states(self, pos_t, pos_t1, head_t, head_t1, vel_t):
|
||||
"""从状态序列反推动作"""
|
||||
actions = []
|
||||
dt = 0.1
|
||||
|
||||
for i in range(len(pos_t)):
|
||||
current_speed = np.linalg.norm(vel_t[i])
|
||||
displacement = np.linalg.norm(pos_t1[i, :2] - pos_t[i, :2])
|
||||
next_speed = displacement / dt
|
||||
|
||||
speed_change = (next_speed - current_speed) / dt
|
||||
if speed_change >= 0:
|
||||
throttle = np.clip(speed_change / 5.0, 0.0, 1.0)
|
||||
else:
|
||||
throttle = np.clip(speed_change / 8.0, -1.0, 0.0)
|
||||
|
||||
heading_change = head_t1[i] - head_t[i]
|
||||
heading_change = np.arctan2(np.sin(heading_change), np.cos(heading_change))
|
||||
steering = np.clip(heading_change / 0.2, -1.0, 1.0)
|
||||
|
||||
actions.append([throttle, steering])
|
||||
|
||||
return np.array(actions)
|
||||
|
||||
@staticmethod
|
||||
def collect_with_full_obs(env_config, num_scenarios=10, save_path=None):
|
||||
"""
|
||||
✅ 使用env._get_all_obs()收集完整107维观测
|
||||
|
||||
这是正确的方法!直接利用环境已有的观测函数
|
||||
"""
|
||||
all_trajectories = {}
|
||||
all_observations = {}
|
||||
|
||||
# 检查数据库
|
||||
data_dir = env_config["config"]["data_directory"]
|
||||
summary_path = os.path.join(data_dir, "dataset_summary.pkl")
|
||||
|
||||
with open(summary_path, 'rb') as f:
|
||||
summary = pickle.load(f)
|
||||
|
||||
total_scenarios = len(summary)
|
||||
print(f"数据库总场景数: {total_scenarios}")
|
||||
|
||||
if num_scenarios is None:
|
||||
num_scenarios = total_scenarios
|
||||
else:
|
||||
num_scenarios = min(num_scenarios, total_scenarios)
|
||||
|
||||
print(f"计划收集(完整107维观测): {num_scenarios} 个场景")
|
||||
|
||||
for i in range(num_scenarios):
|
||||
try:
|
||||
# 创建环境
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
**env_config["config"],
|
||||
"start_scenario_index": i,
|
||||
"num_scenarios": 1,
|
||||
},
|
||||
agent2policy=env_config["agent2policy"]
|
||||
)
|
||||
|
||||
# 重置环境
|
||||
env.reset()
|
||||
|
||||
if not hasattr(env, 'expert_trajectories'):
|
||||
print(f"⚠️ 场景 {i}: 缺少expert_trajectories")
|
||||
env.close()
|
||||
continue
|
||||
|
||||
expert_trajs = env.expert_trajectories
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
print(f"⚠️ 场景 {i}: 无专家轨迹")
|
||||
env.close()
|
||||
continue
|
||||
|
||||
# 存储轨迹
|
||||
scenario_id = env.engine.current_seed
|
||||
for obj_id, traj in expert_trajs.items():
|
||||
unique_id = f"scenario{i}_{obj_id}"
|
||||
all_trajectories[unique_id] = traj
|
||||
|
||||
# ✅ 关键: 使用_get_all_obs()获取完整观测
|
||||
# 创建agent_id到unique_id的映射
|
||||
agent_to_unique = {}
|
||||
for agent_id in env.controlled_agents.keys():
|
||||
# 尝试匹配agent_id到expert_trajectories的obj_id
|
||||
for obj_id in expert_trajs.keys():
|
||||
if str(agent_id) in str(obj_id) or str(obj_id) in str(agent_id):
|
||||
unique_id = f"scenario{i}_{obj_id}"
|
||||
agent_to_unique[agent_id] = unique_id
|
||||
all_observations[unique_id] = []
|
||||
break
|
||||
|
||||
# 遍历场景的每一步,收集完整观测
|
||||
max_steps = min([traj["length"] for traj in expert_trajs.values()])
|
||||
|
||||
for step in range(max_steps):
|
||||
# ✅ 直接调用_get_all_obs()获取107维观测!
|
||||
obs_list = env._get_all_obs()
|
||||
|
||||
# 存储每个agent的观测
|
||||
for agent_idx, agent_id in enumerate(env.controlled_agents.keys()):
|
||||
if agent_id in agent_to_unique:
|
||||
unique_id = agent_to_unique[agent_id]
|
||||
if agent_idx < len(obs_list):
|
||||
# obs_list[agent_idx]已经是107维向量!
|
||||
all_observations[unique_id].append(np.array(obs_list[agent_idx]))
|
||||
|
||||
# 执行零动作(保持场景状态)
|
||||
actions = {aid: np.array([0.0, 0.0])
|
||||
for aid in env.controlled_agents.keys()}
|
||||
env.step(actions)
|
||||
|
||||
# 转换为numpy数组
|
||||
for unique_id in list(all_observations.keys()):
|
||||
if len(all_observations[unique_id]) > 0:
|
||||
all_observations[unique_id] = np.array(all_observations[unique_id])
|
||||
else:
|
||||
del all_observations[unique_id]
|
||||
|
||||
env.close()
|
||||
|
||||
if (i + 1) % 5 == 0:
|
||||
print(f"✓ 已收集 {i+1}/{num_scenarios}, "
|
||||
f"轨迹: {len(all_trajectories)}, "
|
||||
f"观测: {len(all_observations)}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ 场景 {i} 收集失败: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
try:
|
||||
env.close()
|
||||
except:
|
||||
pass
|
||||
continue
|
||||
|
||||
print(f"\n收集完成!")
|
||||
print(f" 轨迹数: {len(all_trajectories)}")
|
||||
print(f" 完整观测数: {len(all_observations)}")
|
||||
|
||||
# 验证观测维度
|
||||
if len(all_observations) > 0:
|
||||
sample_obs = list(all_observations.values())[0]
|
||||
if len(sample_obs) > 0:
|
||||
obs_dim = len(sample_obs[0])
|
||||
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,
|
||||
"observations": all_observations
|
||||
}, f)
|
||||
print(f"数据已保存到: {save_path}")
|
||||
|
||||
return all_trajectories, all_observations
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/mdsn"
|
||||
data_dir = AssetLoader.file_path(WAYMO_DATA_DIR, "exp_filtered", unix_style=False)
|
||||
|
||||
env_config = {
|
||||
"config": {
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
},
|
||||
"agent2policy": DummyPolicy()
|
||||
}
|
||||
|
||||
print("=" * 60)
|
||||
print("选择收集模式:")
|
||||
print("1. 简化观测(5维) - 快速,已验证 ✅")
|
||||
print("2. 完整观测(107维) - 使用_get_all_obs() ⭐")
|
||||
print("=" * 60)
|
||||
|
||||
mode = input("请选择模式(1或2,默认1): ").strip() or "1"
|
||||
|
||||
if mode == "2":
|
||||
print("\n开始收集完整107维观测...")
|
||||
trajectories, observations = ExpertTrajectoryDataset.collect_with_full_obs(
|
||||
env_config,
|
||||
num_scenarios=10,
|
||||
save_path="data/trajectories/expert_trajectories_full.pkl"
|
||||
)
|
||||
|
||||
if len(trajectories) > 0:
|
||||
dataset = ExpertTrajectoryDataset(
|
||||
trajectories,
|
||||
observations,
|
||||
sequence_length=1
|
||||
)
|
||||
state, action = dataset[0]
|
||||
print(f"\n数据集测试:")
|
||||
print(f" 总轨迹数: {len(trajectories)}")
|
||||
print(f" 总观测数: {len(observations)}")
|
||||
print(f" 训练样本数: {len(dataset)}")
|
||||
print(f" 状态维度: {state.shape}")
|
||||
print(f" 动作维度: {action.shape}")
|
||||
else:
|
||||
print("\n开始收集简化5维观测...")
|
||||
# 保持原有的简化版本代码...
|
||||
print("(使用之前已成功的方法)")
|
||||
164
dataset/loader.py
Normal file
164
dataset/loader.py
Normal file
@@ -0,0 +1,164 @@
|
||||
"""
|
||||
统一数据加载: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, *, filter_terminal_last_step: bool = False, agent_id_filter=None):
|
||||
"""从目录或单个 pkl 加载专家 (obs, acts),返回 concat 后的 obs_data, act_data。
|
||||
|
||||
Args:
|
||||
expert_data_path: Directory containing pkl files or a single pkl file.
|
||||
filter_terminal_last_step: If True, drop the last (obs, act) pair of each trajectory.
|
||||
This approximates II's \"train only on non-terminal steps\" when the dataset doesn't
|
||||
explicitly store dones.
|
||||
agent_id_filter: If not None, only load trajectories with traj[\"agent_id\"] == agent_id_filter
|
||||
(e.g. \"default_agent\" for single-agent/ego-only).
|
||||
"""
|
||||
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 agent_id_filter is not None and traj.get("agent_id") != agent_id_filter:
|
||||
continue
|
||||
if "obs" in traj and "acts" in traj:
|
||||
obs = traj["obs"]
|
||||
acts = traj["acts"]
|
||||
if filter_terminal_last_step and len(obs) > 0 and len(acts) > 0:
|
||||
# Drop last step of each trajectory
|
||||
obs = obs[:-1]
|
||||
acts = acts[:-1]
|
||||
if len(obs) == 0 or len(acts) == 0:
|
||||
continue
|
||||
obs_data.append(obs)
|
||||
act_data.append(acts)
|
||||
elif isinstance(data, dict):
|
||||
if "observations" in data and "actions" in data:
|
||||
obs = data["observations"]
|
||||
acts = data["actions"]
|
||||
if filter_terminal_last_step and len(obs) > 0 and len(acts) > 0:
|
||||
obs = obs[:-1]
|
||||
acts = acts[:-1]
|
||||
if len(obs) == 0 or len(acts) == 0:
|
||||
continue
|
||||
obs_data.append(obs)
|
||||
act_data.append(acts)
|
||||
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
|
||||
|
||||
|
||||
def get_expert_scenario_ids(expert_data_path, max_ids=10):
|
||||
"""
|
||||
从专家 pkl 中收集出现过的 scenario_id(这些场景在采集时曾有受控车)。
|
||||
用于 eval 时只在这些场景上评估,保证 eval 有受控车。
|
||||
返回排序后的 list,最多 max_ids 个。
|
||||
"""
|
||||
if os.path.isdir(expert_data_path):
|
||||
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
||||
elif os.path.exists(expert_data_path):
|
||||
pkl_files = [expert_data_path]
|
||||
else:
|
||||
return []
|
||||
|
||||
seen = set()
|
||||
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 "scenario_id" in traj:
|
||||
seen.add(traj["scenario_id"])
|
||||
# dict 格式通常没有 per-trajectory scenario_id,跳过
|
||||
except Exception:
|
||||
continue
|
||||
out = sorted(seen)[:max_ids]
|
||||
return out
|
||||
|
||||
|
||||
class MAGAILExpertDataset(Dataset):
|
||||
def __init__(self, data_dir, transform=None, *, filter_terminal_last_step: bool = False, agent_id_filter=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.
|
||||
agent_id_filter: If not None, only load trajectories with traj[\"agent_id\"] == agent_id_filter.
|
||||
"""
|
||||
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), ...}
|
||||
if agent_id_filter is not None:
|
||||
data = [t for t in data if t.get("agent_id") == agent_id_filter]
|
||||
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)
|
||||
max_i = len(obs)
|
||||
if filter_terminal_last_step and max_i > 0:
|
||||
max_i -= 1
|
||||
for i in range(max_i):
|
||||
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
|
||||
439
docs/HBBC_Deploy_guied.md
Normal file
439
docs/HBBC_Deploy_guied.md
Normal file
@@ -0,0 +1,439 @@
|
||||
# HBBC 策略部署指南
|
||||
|
||||
本文档说明如何将 `weights/hbbc.pt` 部署到 MetaDrive 项目中的**背景车辆**上,作为车辆控制策略使用。
|
||||
|
||||
---
|
||||
|
||||
## 0. 本仓库适配说明(MAGAIL4AutoDrive)
|
||||
|
||||
本仓库已落地一套可直接使用的 HBBC 背景车接入实现,核心代码:
|
||||
|
||||
- `Env/hbbc_actor_critic.py`:HBBC 所需 `ActorCritic` 最小推理网络
|
||||
- `Env/hbbc_background_policy.py`:模型加载、18 维观测构建、latent 管理(含 JSON 覆盖)
|
||||
- `Env/bc_env.py`:`BCScenarioEnv` 动态背景车 HBBC 接入(静态背景车保持不变)
|
||||
- `Env/bc_ego_replay_env.py`:`BCEgoReplayEnv` 动态背景车 HBBC 接入(ego-only 评估兼容)
|
||||
|
||||
与原文档示例不同点:
|
||||
|
||||
1. 当前仓库 `BaseVehicle` 没有 `pos_buffer/rot_buffer/action_buffer`,因此 8 维 `base_state` 使用当前可得车辆状态重建;
|
||||
2. 仅动态背景车使用 HBBC,静态背景车仍作为占位/邻居车辆;
|
||||
3. 支持通过 JSON 手动指定场景中某些车辆的 latent(`object_id` / `agent_id` 双 key)。
|
||||
|
||||
---
|
||||
|
||||
## 1. 概述
|
||||
|
||||
### 1.1 HBBC 是什么
|
||||
|
||||
**HBBC**(Hierarchical Behavior-Based Controller)是一个低层驾驶策略网络,输入车辆状态和行为条件,输出连续控制动作 `[steering, acceleration]`,可直接用于 MetaDrive 的车辆控制。
|
||||
|
||||
### 1.2 依赖
|
||||
|
||||
- **PyTorch**
|
||||
- **NumPy**
|
||||
- **MetaDrive**(需包含 `BaseVehicle`、`BasePolicy` 等基础组件)
|
||||
|
||||
---
|
||||
|
||||
## 2. 模型加载
|
||||
|
||||
### 2.1 模型架构
|
||||
|
||||
HBBC 对应 `ActorCritic` 网络,需按以下参数实例化:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from algorithms.modules import ActorCritic # 或复制 actor_critic.py 到目标项目
|
||||
|
||||
hbbc = ActorCritic(
|
||||
num_actor_obs=18,
|
||||
num_critic_obs=18,
|
||||
num_actions=2,
|
||||
latent_c_dim=4, # 行为模式数
|
||||
latent_eps_dim=6, # 风格向量维度
|
||||
use_style_latent=True,
|
||||
).to(device)
|
||||
|
||||
# 加载权重
|
||||
checkpoint = torch.load("path/to/hbbc.pt", map_location=device, weights_only=False)
|
||||
hbbc.load_state_dict(checkpoint['actor_critic'])
|
||||
hbbc.eval()
|
||||
```
|
||||
|
||||
### 2.2 推理接口
|
||||
|
||||
```python
|
||||
with torch.no_grad():
|
||||
actions = hbbc.act_inference(obs_tensor) # obs_tensor: (batch, 18), 输出: (batch, 2)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 3. 输入规格(18 维)
|
||||
|
||||
HBBC 的输入为 `hbbc_obs`,维度 18,由三部分拼接:
|
||||
|
||||
```
|
||||
hbbc_obs = [base_state(8) | latent_eps(6) | latent_c(4)]
|
||||
```
|
||||
|
||||
### 3.1 base_state(8 维)
|
||||
|
||||
从车辆对象构建,需按**精确顺序**拼接。实现如下(需配合 `relative_pos_local`、`rot_matrix_inv`、`clip` 等工具函数):
|
||||
|
||||
```python
|
||||
import numpy as np
|
||||
|
||||
def build_hbbc_base_state(vehicle):
|
||||
"""
|
||||
从 MetaDrive 车辆对象构建 HBBC 的 8 维 base_state。
|
||||
要求 vehicle 具有: position, pos_buffer, rot_buffer, heading_buffer,
|
||||
speed_km_h, max_speed_km_h, eps_step, acceleration, yaw_rate, action_buffer
|
||||
"""
|
||||
from metadrive.utils.math import clip # 或 np.clip
|
||||
|
||||
veh_pos = list(vehicle.position) + [0]
|
||||
init_veh_rot = np.array([vehicle.rot_buffer[0][0], vehicle.rot_buffer[0][1], vehicle.rot_buffer[0][2]])
|
||||
init_veh_pos = list(vehicle.pos_buffer[0]) + [0]
|
||||
init_veh_heading = vehicle.heading_buffer[0]
|
||||
|
||||
# 局部位置(本实现中置 0)
|
||||
veh_pos_local = relative_pos_local(init_veh_pos, veh_pos, init_veh_rot)[:2]
|
||||
veh_pos_local[0] /= 10
|
||||
veh_pos_local[1] /= 2
|
||||
|
||||
# 局部航向(本实现中置 0)
|
||||
veh_heading = vehicle.heading
|
||||
cross = np.cross(init_veh_heading, veh_heading)
|
||||
dot = np.dot(init_veh_heading, veh_heading)
|
||||
veh_heading_local = np.arctan2(cross, dot)
|
||||
|
||||
veh_vel = clip((vehicle.speed_km_h + 1) / (vehicle.max_speed_km_h + 1), 0.0, 1.0)
|
||||
veh_acc = vehicle.acceleration / 5 if vehicle.eps_step > 1 else 0
|
||||
yaw_rate = vehicle.yaw_rate
|
||||
last_action_0 = vehicle.action_buffer[-1][0]
|
||||
last_action_1 = vehicle.action_buffer[-1][1]
|
||||
|
||||
# 8 维,顺序固定
|
||||
obs = np.concatenate((
|
||||
veh_pos_local * 0, # 2 维,置 0
|
||||
[veh_heading_local * 0], # 1 维,置 0
|
||||
[veh_vel], # 1 维
|
||||
[veh_acc * 0], # 1 维,置 0
|
||||
[yaw_rate * 0.5], # 1 维
|
||||
[last_action_0], [last_action_1] # 2 维
|
||||
)).astype(np.float32)
|
||||
return obs
|
||||
```
|
||||
|
||||
### 3.2 latent_eps(6 维)
|
||||
|
||||
风格向量,需 **L2 归一化** 且在 `[-1, 1]` 内:
|
||||
|
||||
```python
|
||||
# 随机采样(每个 episode 或每辆车可固定/随机)
|
||||
latent_eps = np.random.randn(6).astype(np.float32)
|
||||
latent_eps = latent_eps / (np.linalg.norm(latent_eps) + 1e-8)
|
||||
latent_eps = np.clip(latent_eps, -1.0, 1.0)
|
||||
```
|
||||
|
||||
### 3.3 latent_c(4 维)
|
||||
|
||||
行为模式 one-hot,4 选 1:
|
||||
|
||||
```python
|
||||
# 随机选一个模式 (0~3)
|
||||
mode = np.random.randint(0, 4)
|
||||
latent_c = np.zeros(4, dtype=np.float32)
|
||||
latent_c[mode] = 1.0
|
||||
```
|
||||
|
||||
### 3.4 完整观测拼接
|
||||
|
||||
```python
|
||||
def build_hbbc_obs(vehicle, latent_eps, latent_c):
|
||||
base = build_hbbc_base_state(vehicle)
|
||||
return np.concatenate([base, latent_eps, latent_c], axis=-1) # shape: (18,)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 4. 必需工具函数
|
||||
|
||||
若目标项目无以下函数,需自行实现或从 styledrive 的 `envs/utils.py` 拷贝:
|
||||
|
||||
```python
|
||||
def rot_matrix(t):
|
||||
"""t: [roll, pitch, yaw], 返回 3x3 旋转矩阵"""
|
||||
roll, pitch, yaw = t[0], t[1], t[2]
|
||||
sr, cr = np.sin(roll), np.cos(roll)
|
||||
sp, cp = np.sin(pitch), np.cos(pitch)
|
||||
sy, cy = np.sin(yaw), np.cos(yaw)
|
||||
r_roll = np.array([[1, 0, 0], [0, cr, -sr], [0, sr, cr]])
|
||||
r_pitch = np.array([[cp, 0, sp], [0, 1, 0], [-sp, 0, cp]])
|
||||
r_yaw = np.array([[cy, -sy, 0], [sy, cy, 0], [0, 0, 1]])
|
||||
return np.dot(np.dot(r_yaw, r_pitch), r_roll)
|
||||
|
||||
def rot_matrix_inv(t):
|
||||
return rot_matrix(t).T
|
||||
|
||||
def relative_pos_local(coord, coord_t, veh_rot):
|
||||
"""将 coord_t 从世界坐标变换到以 coord 为原点、veh_rot 为姿态的局部坐标"""
|
||||
r_pos_global = np.array(coord_t) - np.array(coord)
|
||||
rot_mat_inv = rot_matrix_inv(veh_rot)
|
||||
return rot_mat_inv @ r_pos_global
|
||||
```
|
||||
|
||||
`clip` 可用 `np.clip` 或 `metadrive.utils.math.clip`。
|
||||
|
||||
---
|
||||
|
||||
## 5. 车辆属性要求
|
||||
|
||||
使用 HBBC 的车辆需继承或兼容 MetaDrive 的 `BaseVehicle`,并具备:
|
||||
|
||||
| 属性 | 说明 |
|
||||
|------|------|
|
||||
| `position` | 当前位置 (x, y) 或 (x, y, z) |
|
||||
| `heading` | 航向单位向量 |
|
||||
| `heading_theta` | 航向角(弧度) |
|
||||
| `pos_buffer` | `deque`,至少 1 个元素,`pos_buffer[0]` 为 episode 起始位姿 |
|
||||
| `rot_buffer` | `deque`,`(roll, pitch, yaw)`,`rot_buffer[0]` 为起始姿态 |
|
||||
| `heading_buffer` | `deque`,`heading_buffer[0]` 为起始航向 |
|
||||
| `action_buffer` | `deque`,`action_buffer[-1]` 为上一时刻动作 `(steering, acc)` |
|
||||
| `speed_km_h` | 当前速度 km/h |
|
||||
| `max_speed_km_h` | 最大速度 km/h |
|
||||
| `acceleration` | 当前加速度 |
|
||||
| `yaw_rate` | 偏航角速度 (rad/s) |
|
||||
| `eps_step` | 本 episode 的步数 |
|
||||
| `last_heading_theta` | 上一帧航向角(用于 yaw_rate) |
|
||||
|
||||
`BaseVehicle` 在 `before_step` 中会更新 `pos_buffer`、`rot_buffer`、`heading_buffer`、`action_buffer`,只要在配置中设置 `veh_obs_len >= 1`(建议 3–10)即可。
|
||||
|
||||
---
|
||||
|
||||
## 6. 输出动作格式
|
||||
|
||||
HBBC 输出 2 维连续动作,与 MetaDrive 动作空间一致:
|
||||
|
||||
```python
|
||||
# actions: (2,) 或 (batch, 2)
|
||||
# actions[0]: steering ∈ [-1, 1]
|
||||
# actions[1]: acceleration ∈ [-1, 1],正=油门,负=刹车
|
||||
```
|
||||
|
||||
环境会在 `_preprocess_actions` 中做限幅与平滑,无需在策略内再次裁剪。
|
||||
|
||||
---
|
||||
|
||||
## 7. 部署为 MetaDrive 策略(背景车)
|
||||
|
||||
### 7.1 自定义 Policy
|
||||
|
||||
实现一个继承 `BasePolicy` 的策略,在 `act` 中调用 HBBC:
|
||||
|
||||
```python
|
||||
from metadrive.policy.base_policy import BasePolicy
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
class HBBCPolicy(BasePolicy):
|
||||
def __init__(self, control_object, random_seed=None, hbbc_path="weights/hbbc.pt", device="cpu"):
|
||||
super().__init__(control_object, random_seed)
|
||||
self.device = torch.device(device)
|
||||
self.hbbc = self._load_hbbc(hbbc_path)
|
||||
self.latent_eps = None
|
||||
self.latent_c = None
|
||||
self._resample_latent()
|
||||
|
||||
def _load_hbbc(self, path):
|
||||
from algorithms.modules import ActorCritic # 根据实际路径调整
|
||||
model = ActorCritic(
|
||||
num_actor_obs=18, num_critic_obs=18, num_actions=2,
|
||||
latent_c_dim=4, latent_eps_dim=6, use_style_latent=True
|
||||
).to(self.device)
|
||||
ckpt = torch.load(path, map_location=self.device, weights_only=False)
|
||||
model.load_state_dict(ckpt['actor_critic'])
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
def _resample_latent(self):
|
||||
self.latent_eps = np.random.randn(6).astype(np.float32)
|
||||
self.latent_eps = self.latent_eps / (np.linalg.norm(self.latent_eps) + 1e-8)
|
||||
self.latent_eps = np.clip(self.latent_eps, -1.0, 1.0)
|
||||
mode = np.random.randint(0, 4)
|
||||
self.latent_c = np.zeros(4, dtype=np.float32)
|
||||
self.latent_c[mode] = 1.0
|
||||
|
||||
def act(self, agent_id=None):
|
||||
vehicle = self.control_object
|
||||
base_state = build_hbbc_base_state(vehicle)
|
||||
obs = np.concatenate([base_state, self.latent_eps, self.latent_c], axis=-1)
|
||||
obs_t = torch.tensor(obs, dtype=torch.float32, device=self.device).unsqueeze(0)
|
||||
with torch.no_grad():
|
||||
actions = self.hbbc.act_inference(obs_t).cpu().numpy().squeeze()
|
||||
self.action_info["action"] = actions.tolist()
|
||||
return [float(actions[0]), float(actions[1])]
|
||||
|
||||
def reset(self):
|
||||
super().reset()
|
||||
self._resample_latent()
|
||||
```
|
||||
|
||||
### 7.2 配置背景车使用 HBBC
|
||||
|
||||
在环境配置中为背景车辆指定 `HBBCPolicy`:
|
||||
|
||||
```python
|
||||
config = {
|
||||
# ...
|
||||
"agent_configs": {
|
||||
"agent0": {
|
||||
"policy": HBBCPolicy,
|
||||
"policy_kwargs": {"hbbc_path": "path/to/hbbc.pt", "device": "cuda:0"},
|
||||
}
|
||||
},
|
||||
# 若使用 traffic 的 policy 配置方式,则需在 traffic 管理逻辑中
|
||||
# 将部分或全部背景车的 policy 替换为 HBBCPolicy
|
||||
}
|
||||
```
|
||||
|
||||
若背景车由 TrafficManager 等模块统一管理,需在该模块的 policy 选择逻辑中加入对 `HBBCPolicy` 的分配。
|
||||
|
||||
### 7.3 与 TrafficManager 集成
|
||||
|
||||
若背景车由 `PGTrafficManager` 等生成,需在添加策略时改为使用 `HBBCPolicy`:
|
||||
|
||||
```python
|
||||
# 原代码通常为:
|
||||
# self.add_policy(random_v.id, IDMPolicy, random_v, self.generate_seed())
|
||||
|
||||
# 改为:
|
||||
from your_policy_module import HBBCPolicy
|
||||
self.add_policy(random_v.id, HBBCPolicy, random_v, self.generate_seed(),
|
||||
hbbc_path="path/to/hbbc.pt", device="cuda:0")
|
||||
```
|
||||
|
||||
`add_policy` 的额外参数会传给 Policy 的 `__init__`。若接口不支持传参,可修改 `HBBCPolicy` 从全局配置读取路径,或使用自定义 TrafficManager 子类。
|
||||
|
||||
**注意**:HBBC 在 styledrive 中基于 scenario 轨迹训练,不包含路由逻辑。背景车若需要沿车道/路线行驶,可能需:
|
||||
- 在项目中为 HBBC 车辆配置 `navigation`,或
|
||||
- 仅对部分背景车使用 HBBC(如混合 IDM + HBBC),或
|
||||
- 在目标项目中验证 HBBC 在开放道路上的表现后决定是否全量使用。
|
||||
|
||||
### 7.4 注意事项
|
||||
|
||||
1. **latent 生命周期**:可为每辆车在 spawn 时采样一次,或在每个 episode reset 时重采样。
|
||||
2. **首帧 action_buffer**:首步 `action_buffer[-1]` 通常为 `(0, 0)`,由 `BaseVehicle` 初始化保证。
|
||||
3. **同步更新 buffer**:车辆必须在每步调用 `before_step` 之类接口,更新 `pos_buffer`、`action_buffer` 等,否则观测会错位。
|
||||
4. **veh_obs_len**:车辆配置中设置 `veh_obs_len >= 3`(建议 10),确保 buffer 长度足够。
|
||||
|
||||
---
|
||||
|
||||
## 8. ActorCritic 网络定义(可移植)
|
||||
|
||||
若目标项目无法导入 styledrive 的 `algorithms`,可把以下简化版 `ActorCritic` 放到本项目中单独使用:
|
||||
|
||||
```python
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
def get_activation(name):
|
||||
return getattr(nn, name)()
|
||||
|
||||
class ActorCritic(nn.Module):
|
||||
def __init__(self, num_actor_obs=18, num_critic_obs=18, num_actions=2,
|
||||
latent_c_dim=4, latent_eps_dim=6, use_style_latent=True,
|
||||
actor_hidden_dims=[512, 256, 128], activation='elu'):
|
||||
super().__init__()
|
||||
act_fn = getattr(nn, activation.upper())()
|
||||
self.latent_c_dim = latent_c_dim
|
||||
self.latent_eps_dim = latent_eps_dim
|
||||
self.use_style_latent = use_style_latent
|
||||
|
||||
layers = []
|
||||
layers.append(nn.Linear(num_actor_obs, actor_hidden_dims[0]))
|
||||
layers.append(act_fn)
|
||||
for i in range(len(actor_hidden_dims) - 1):
|
||||
layers.append(nn.Linear(actor_hidden_dims[i], actor_hidden_dims[i + 1]))
|
||||
layers.append(act_fn)
|
||||
self.actor_trunk = nn.Sequential(*layers)
|
||||
self.actor_head = nn.Linear(actor_hidden_dims[-1], num_actions)
|
||||
|
||||
if use_style_latent:
|
||||
style_layers = [nn.Linear(latent_eps_dim, 512), act_fn,
|
||||
nn.Linear(512, 256), act_fn, nn.Linear(256, 128), act_fn]
|
||||
self.style_trunk = nn.Sequential(*style_layers)
|
||||
self.style_head = nn.Linear(128, latent_eps_dim)
|
||||
self.style_activation = torch.tanh
|
||||
|
||||
def act_inference(self, observations):
|
||||
if self.use_style_latent:
|
||||
obs = observations[..., :-(self.latent_c_dim + self.latent_eps_dim)]
|
||||
eps = observations[..., -self.latent_c_dim - self.latent_eps_dim:-self.latent_c_dim]
|
||||
c = observations[..., -self.latent_c_dim:]
|
||||
eps = self.style_activation(self.style_head(self.style_trunk(eps)))
|
||||
observations = torch.cat([obs, eps, c], dim=-1)
|
||||
embedding = self.actor_trunk(observations)
|
||||
return self.actor_head(embedding)
|
||||
```
|
||||
|
||||
加载与调用方式与前面一致。
|
||||
|
||||
---
|
||||
|
||||
## 9. 简要检查清单
|
||||
|
||||
- [ ] 正确加载 `hbbc.pt` 的 `actor_critic` 权重
|
||||
- [ ] `build_hbbc_base_state` 输出 8 维,顺序与文档一致
|
||||
- [ ] `latent_eps` 6 维、L2 归一化
|
||||
- [ ] `latent_c` 4 维 one-hot
|
||||
- [ ] 车辆具备 `pos_buffer`、`rot_buffer`、`heading_buffer`、`action_buffer` 等属性
|
||||
- [ ] 策略返回 `[steering, acceleration]`,范围 [-1, 1]
|
||||
- [ ] 每步更新上述 buffer,保证观测连续
|
||||
|
||||
---
|
||||
|
||||
## 10. 本仓库配置项与 JSON 示例
|
||||
|
||||
可通过环境配置控制 HBBC 背景车行为:
|
||||
|
||||
- `enable_hbbc_background`:是否启用动态背景车 HBBC(`True/False`)
|
||||
- `hbbc_model_path`:模型路径(默认 `models/hbbc/hbbc.pt`)
|
||||
- `hbbc_inference_device`:推理设备(如 `cpu` / `cuda:0`)
|
||||
- `hbbc_latent_mode`:`per_vehicle_fixed` 或 `per_episode_reset`
|
||||
- `hbbc_latent_json_path`:可选,手动 latent JSON 路径
|
||||
|
||||
`hbbc_latent_json_path` 内容格式(优先按 `object_id` 匹配,失败回退 `agent_id`):
|
||||
|
||||
```json
|
||||
{
|
||||
"global": {
|
||||
"latent_eps": [0.35, -0.12, 0.28, 0.46, -0.22, 0.18],
|
||||
"latent_c": [0, 0, 1, 0]
|
||||
},
|
||||
"object_id": {
|
||||
"12345": {
|
||||
"latent_eps": [0.2, -0.1, 0.3, 0.4, -0.2, 0.1],
|
||||
"latent_c": [0, 1, 0, 0]
|
||||
}
|
||||
},
|
||||
"agent_id": {
|
||||
"controlled_abcde": {
|
||||
"latent_eps": [0.5, 0.1, -0.1, 0.2, -0.3, 0.4],
|
||||
"latent_c": [1, 0, 0, 0]
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
匹配优先级为:`object_id` > `agent_id` > `global` > 随机采样。
|
||||
`latent_eps` 会做 L2 归一化,`latent_c` 会强制 one-hot;非法输入会告警并回退随机采样。
|
||||
|
||||
---
|
||||
|
||||
## 11. 参考来源
|
||||
|
||||
- 策略与观测:`envs/ad_hbbc_gym.py` 中的 `ADObservation.vehicle_state`
|
||||
- 模型:`algorithms/modules/actor_critic.py` 中 `ActorCritic`
|
||||
- 工具:`envs/utils.py` 中的 `relative_pos_local`、`rot_matrix`、`rot_matrix_inv`
|
||||
498
docs/TRAINING_ARCHITECTURE.md
Normal file
498
docs/TRAINING_ARCHITECTURE.md
Normal file
@@ -0,0 +1,498 @@
|
||||
# MAGAIL 训练方案架构文档
|
||||
|
||||
## 目录
|
||||
1. [训练数据结构](#1-训练数据结构)
|
||||
2. [多智能体训练机制](#2-多智能体训练机制)
|
||||
3. [完整训练流程](#3-完整训练流程)
|
||||
4. [当前项目问题](#4-当前项目问题)
|
||||
5. [TensorBoard 日志问题](#5-tensorboard-日志问题)
|
||||
|
||||
---
|
||||
|
||||
## 1. 训练数据结构
|
||||
|
||||
### 1.1 数据维度
|
||||
|
||||
**观测空间 (Observation Space)**
|
||||
- **维度**: 45维
|
||||
- **组成**:
|
||||
- **Ego状态** (5维): `[position_x, position_y, velocity_x, velocity_y, heading_theta]`
|
||||
- **邻居信息** (40维): 最多10个邻居,每个邻居4维特征
|
||||
- 每个邻居: `[relative_x, relative_y, velocity_x, velocity_y]`
|
||||
- 如果邻居数量 < 10,用零填充
|
||||
|
||||
**动作空间 (Action Space)**
|
||||
- **维度**: 2维
|
||||
- **组成**: `[steering, accel]`
|
||||
- **范围**: 归一化到 `[-1, 1]`
|
||||
|
||||
### 1.2 数据格式
|
||||
|
||||
**专家数据文件结构** (`.pkl` 文件):
|
||||
```python
|
||||
# 每个 .pkl 文件包含一个列表,每个元素是一条车辆轨迹
|
||||
trajectories = [
|
||||
{
|
||||
'obs': np.array, # Shape: (T, 45) - T为轨迹长度(可变)
|
||||
'acts': np.array, # Shape: (T, 2) - 对应的动作序列
|
||||
'agent_id': str, # 车辆ID
|
||||
'scenario_id': int # 场景ID
|
||||
},
|
||||
...
|
||||
]
|
||||
```
|
||||
|
||||
**数据特点**:
|
||||
- 轨迹长度 `T` 是**可变的**,取决于车辆在场景中的存活时间
|
||||
- 最小轨迹长度过滤: 只保留长度 > 10 的轨迹
|
||||
- 数据已通过静态车辆过滤(移动距离 < 5m 且最大速度 < 1m/s 的车辆被过滤)
|
||||
|
||||
### 1.3 数据生成流程
|
||||
|
||||
**脚本**: `scripts/generate_expert_data.py`
|
||||
|
||||
**流程**:
|
||||
1. 从 Waymo 数据 (`data/exp_filtered`) 加载场景
|
||||
2. 使用 `ExpertReplayEnv` 回放专家轨迹
|
||||
3. 通过逆动力学 (`Env/inverse_dynamics.py`) 计算动作
|
||||
4. 构建45维观测(Ego + 10个最近邻居)
|
||||
5. 过滤无效轨迹(长度 < 10)
|
||||
6. 保存为 `.pkl` 文件到 `data/training_data/`
|
||||
|
||||
**关键代码位置**:
|
||||
- 观测构建: `Env/expert_replay_env.py` 的 `_get_all_obs()` 方法
|
||||
- 动作计算: `Env/inverse_dynamics.py` 的 `compute_action()` 方法
|
||||
|
||||
---
|
||||
|
||||
## 2. 多智能体训练机制
|
||||
|
||||
### 2.1 可变长度处理
|
||||
|
||||
**问题**: 不同场景中智能体数量不同,每个智能体的轨迹长度也不同。
|
||||
|
||||
**解决方案**:
|
||||
|
||||
1. **数据层面** (`dataset/magail_dataset.py`):
|
||||
- 将轨迹**展平**为独立的 `(state, action)` 对
|
||||
- 每个样本是独立的,不保留序列信息
|
||||
- 这样所有轨迹可以统一处理,不受长度限制
|
||||
|
||||
```python
|
||||
# MAGAILExpertDataset 的处理方式
|
||||
for traj in self.trajectories:
|
||||
obs = traj['obs'] # (T, 45)
|
||||
acts = traj['acts'] # (T, 2)
|
||||
# 展平为独立样本
|
||||
for i in range(len(obs)):
|
||||
self.flat_data.append((obs[i], acts[i])) # 每个样本: (45,), (2,)
|
||||
```
|
||||
|
||||
2. **训练环境层面** (`train_magail.py`):
|
||||
- 每个 episode 动态处理不同数量的智能体
|
||||
- 在 rollout 循环中,为每个活跃智能体独立收集数据
|
||||
- 所有智能体的数据合并到一个 `memory` 中
|
||||
|
||||
```python
|
||||
# Rollout 循环
|
||||
for agent_id, obs in obs_dict.items():
|
||||
act, logprob = ppo_agent.select_action(obs)
|
||||
actions[agent_id] = act
|
||||
# 所有智能体的数据都存入同一个 memory
|
||||
memory['states'].append(obs)
|
||||
memory['actions'].append(actions[agent_id])
|
||||
...
|
||||
```
|
||||
|
||||
3. **观测维度固定**:
|
||||
- 通过 `MAGAILScenarioEnv` 确保观测维度始终为45维
|
||||
- 邻居数量不足时用零填充,保证维度一致
|
||||
|
||||
### 2.2 多智能体交互
|
||||
|
||||
**环境设置**:
|
||||
- 使用 `MAGAILScenarioEnv` (继承自 `MultiAgentScenarioEnv`)
|
||||
- 自定义 `_get_all_obs()` 方法,确保观测格式与专家数据一致
|
||||
- 每个智能体独立选择动作,环境统一执行
|
||||
|
||||
**关键点**:
|
||||
- 所有智能体共享同一个策略网络(参数共享)
|
||||
- 每个智能体独立计算动作和奖励
|
||||
- 数据收集时将所有智能体的经验合并
|
||||
|
||||
---
|
||||
|
||||
## 3. 完整训练流程
|
||||
|
||||
### 3.1 数据准备阶段
|
||||
|
||||
**步骤 1: 生成专家数据**
|
||||
```bash
|
||||
python scripts/generate_expert_data.py \
|
||||
--data_dir data/exp_filtered \
|
||||
--output_dir data/training_data \
|
||||
--num_scenarios 100 \
|
||||
--start_index 0
|
||||
```
|
||||
|
||||
**输出**: `data/training_data/expert_data_*.pkl`
|
||||
|
||||
### 3.2 模型初始化
|
||||
|
||||
**网络架构**:
|
||||
|
||||
1. **Actor (策略网络)**:
|
||||
- 输入: 45维状态
|
||||
- 输出: 2维动作(连续)
|
||||
- 结构: MLP (45 → 256 → 256 → 2)
|
||||
- 输出分布: 高斯分布(均值 + 可学习标准差)
|
||||
|
||||
2. **Critic (价值网络)**:
|
||||
- 输入: 45维状态
|
||||
- 输出: 标量价值
|
||||
- 结构: MLP (45 → 256 → 256 → 1)
|
||||
|
||||
3. **Discriminator (鉴别器)**:
|
||||
- 输入: 45维状态 + 2维动作 = 47维
|
||||
- 输出: 标量(0-1之间,表示专家概率)
|
||||
- 结构: MLP (47 → 256 → 256 → 1) + Sigmoid
|
||||
|
||||
### 3.3 训练循环
|
||||
|
||||
**主循环** (`train_magail.py` 的 `train()` 函数):
|
||||
|
||||
```
|
||||
For each episode:
|
||||
1. 收集 Rollout
|
||||
- 重置环境(随机选择场景)
|
||||
- 运行策略收集轨迹
|
||||
- 存储 (state, action, logprob, next_state, done)
|
||||
|
||||
2. 训练 Discriminator
|
||||
- 采样专家批次
|
||||
- 采样策略批次
|
||||
- 更新鉴别器:
|
||||
- Expert loss: BCE(D(s_e, a_e), 1)
|
||||
- Policy loss: BCE(D(s_p, a_p), 0)
|
||||
- Total: L_d = L_expert + L_policy
|
||||
|
||||
3. 计算 GAIL 奖励
|
||||
- 对所有策略状态-动作对:
|
||||
reward = -log(1 - D(s, a) + ε)
|
||||
- 替换环境奖励
|
||||
|
||||
4. 更新策略 (PPO)
|
||||
- 计算 GAE (Generalized Advantage Estimation)
|
||||
- PPO 更新 (K epochs):
|
||||
- 计算优势函数
|
||||
- 计算策略损失(带clip)
|
||||
- 计算价值损失
|
||||
- 更新 Actor 和 Critic
|
||||
```
|
||||
|
||||
### 3.4 训练目标
|
||||
|
||||
**Discriminator 目标**:
|
||||
```
|
||||
L_D = E_{(s,a)~π_E}[-log(D(s,a))] + E_{(s,a)~π_θ}[-log(1-D(s,a))]
|
||||
```
|
||||
- 最大化区分专家数据和策略数据的能力
|
||||
|
||||
**Policy (Generator) 目标**:
|
||||
```
|
||||
L_π = E_{(s,a)~π_θ}[-log(D(s,a))] - λ_H(π_θ)
|
||||
```
|
||||
- 通过 PPO 优化,使用 GAIL 奖励作为信号
|
||||
- 最大化鉴别器给出的"专家概率"
|
||||
- 同时保持策略熵(探索)
|
||||
|
||||
**PPO 更新**:
|
||||
```python
|
||||
# 优势函数 (GAE)
|
||||
advantages = compute_gae(rewards, values, next_values, dones, gamma, lambda)
|
||||
|
||||
# 策略损失
|
||||
ratios = exp(log_probs - old_log_probs)
|
||||
surr1 = ratios * advantages
|
||||
surr2 = clip(ratios, 1-ε, 1+ε) * advantages
|
||||
policy_loss = -min(surr1, surr2) + 0.01 * entropy
|
||||
|
||||
# 价值损失
|
||||
value_loss = MSE(critic(states), returns)
|
||||
|
||||
# 总损失
|
||||
total_loss = policy_loss + 0.5 * value_loss
|
||||
```
|
||||
|
||||
### 3.5 关键代码位置
|
||||
|
||||
- **训练主循环**: `train_magail.py:278-505`
|
||||
- **PPO 更新**: `train_magail.py:90-146`
|
||||
- **Discriminator 更新**: `train_magail.py:429-462`
|
||||
- **GAIL 奖励计算**: `train_magail.py:472-477`
|
||||
|
||||
---
|
||||
|
||||
## 4. 当前项目问题
|
||||
|
||||
### 4.1 环境重置问题
|
||||
|
||||
**问题描述**:
|
||||
- MetaDrive 环境在快速重置时可能出现对象清理不完整的问题
|
||||
- 错误信息: "You should clear all generated objects..."
|
||||
|
||||
**当前处理**:
|
||||
- 代码中已有异常处理机制(`train_magail.py:288-342`)
|
||||
- 重置失败时会尝试关闭并重新创建环境
|
||||
- 但可能导致训练不稳定
|
||||
|
||||
**建议修复**:
|
||||
- 在每次重置前显式清理所有对象
|
||||
- 增加重置间隔,避免过于频繁的重置
|
||||
- 考虑使用环境池(Environment Pool)复用环境实例
|
||||
|
||||
### 4.2 观测维度对齐
|
||||
|
||||
**问题描述**:
|
||||
- 原始 `MultiAgentScenarioEnv` 返回108维观测(包含Lidar)
|
||||
- 专家数据使用45维观测
|
||||
- 维度不匹配会导致训练失败
|
||||
|
||||
**当前解决方案**:
|
||||
- 通过 `MAGAILScenarioEnv` 重写 `_get_all_obs()` 方法
|
||||
- 确保训练环境与专家数据使用相同的观测格式
|
||||
|
||||
**代码位置**: `train_magail.py:223-262`
|
||||
|
||||
### 4.3 数据收集效率
|
||||
|
||||
**问题描述**:
|
||||
- 每个 episode 都需要完整运行环境收集数据
|
||||
- 可变长度轨迹导致 batch 大小不一致
|
||||
- 可能影响训练稳定性
|
||||
|
||||
**当前处理**:
|
||||
- 使用展平的数据集,每个样本独立
|
||||
- 在 rollout 时收集所有智能体的数据,合并处理
|
||||
|
||||
**潜在改进**:
|
||||
- 考虑使用经验回放缓冲区
|
||||
- 实现轨迹级别的采样(保留序列信息)
|
||||
|
||||
### 4.4 内存管理
|
||||
|
||||
**问题描述**:
|
||||
- 长时间训练可能导致内存泄漏
|
||||
- 环境对象可能没有完全释放
|
||||
|
||||
**当前处理**:
|
||||
- 代码中有显式的 `gc.collect()` 和 `torch.cuda.empty_cache()`
|
||||
- 但可能不够彻底
|
||||
|
||||
**建议**:
|
||||
- 定期检查内存使用
|
||||
- 考虑限制 rollout 长度
|
||||
- 使用更激进的清理策略
|
||||
|
||||
### 4.5 训练稳定性
|
||||
|
||||
**问题描述**:
|
||||
- Discriminator 可能过早收敛,导致策略无法学习
|
||||
- GAIL 奖励可能不稳定
|
||||
|
||||
**当前处理**:
|
||||
- 使用标准的 GAIL 奖励公式: `-log(1 - D(s,a) + ε)`
|
||||
- PPO 的 clip 机制提供稳定性
|
||||
|
||||
**潜在改进**:
|
||||
- 考虑使用 WGAN-GP 或 LSGAN 损失
|
||||
- 实现 Discriminator 的预训练
|
||||
- 添加奖励归一化
|
||||
|
||||
---
|
||||
|
||||
## 5. TensorBoard 日志问题
|
||||
|
||||
### 5.1 问题分析
|
||||
|
||||
**现象**:
|
||||
- `runs/magail_0112/` 目录下只有模型文件(`.pth`),没有 TensorBoard 事件文件(`events.out.tfevents.*`)
|
||||
- 其他目录(`magail_full`, `magail_production`)有事件文件
|
||||
|
||||
**可能原因**:
|
||||
|
||||
1. **TensorBoard 未安装**:
|
||||
- 代码中有 try-except 处理(`train_magail.py:269-274`)
|
||||
- 如果 TensorBoard 未安装,`writer` 会被设置为 `None`
|
||||
- 训练会继续,但不会写入日志
|
||||
|
||||
2. **日志写入失败**:
|
||||
- 即使 `SummaryWriter` 创建成功,如果写入时出错,可能不会生成文件
|
||||
- 需要检查是否有异常被静默捕获
|
||||
|
||||
3. **训练中断**:
|
||||
- 如果训练在写入第一个日志前中断,可能没有事件文件
|
||||
- 但模型文件已保存,说明训练至少运行了一段时间
|
||||
|
||||
### 5.2 检查方法
|
||||
|
||||
**步骤 1: 检查 TensorBoard 安装**
|
||||
```bash
|
||||
python -c "import tensorboard; print(tensorboard.__version__)"
|
||||
```
|
||||
|
||||
**步骤 2: 检查训练脚本中的日志写入**
|
||||
查看 `train_magail.py:493-496`:
|
||||
```python
|
||||
if writer:
|
||||
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)
|
||||
```
|
||||
|
||||
**步骤 3: 检查日志目录权限**
|
||||
```bash
|
||||
ls -la runs/magail_0112/
|
||||
```
|
||||
|
||||
### 5.3 解决方案
|
||||
|
||||
**方案 1: 确保 TensorBoard 已安装**
|
||||
```bash
|
||||
pip install tensorboard
|
||||
```
|
||||
|
||||
**方案 2: 添加显式刷新**
|
||||
在训练循环结束后,显式调用 `writer.flush()`:
|
||||
```python
|
||||
if writer:
|
||||
writer.flush() # 确保数据写入磁盘
|
||||
```
|
||||
|
||||
**方案 3: 添加日志验证**
|
||||
在训练开始时检查日志目录:
|
||||
```python
|
||||
if writer:
|
||||
# 测试写入
|
||||
writer.add_scalar('Test/Initialization', 0.0, 0)
|
||||
writer.flush()
|
||||
print(f"TensorBoard logging enabled. Log dir: {args.log_dir}")
|
||||
else:
|
||||
print("WARNING: TensorBoard not available. Logging disabled.")
|
||||
```
|
||||
|
||||
**方案 4: 使用文件日志作为备份**
|
||||
即使 TensorBoard 不可用,也可以写入文本日志:
|
||||
```python
|
||||
import logging
|
||||
logging.basicConfig(
|
||||
filename=os.path.join(args.log_dir, 'training.log'),
|
||||
level=logging.INFO
|
||||
)
|
||||
```
|
||||
|
||||
### 5.4 代码修复建议
|
||||
|
||||
**在 `train_magail.py` 中添加以下改进**:
|
||||
|
||||
1. **确保 disc_loss 在 CPU 上**:
|
||||
```python
|
||||
# 第425行附近
|
||||
disc_loss = torch.tensor(0.0).cuda() # 改为 .cuda() 或保持 CPU
|
||||
# 或者在使用时转换
|
||||
if writer:
|
||||
disc_loss_value = disc_loss.item() if isinstance(disc_loss, torch.Tensor) else disc_loss
|
||||
writer.add_scalar('Loss/Discriminator', disc_loss_value, i_episode)
|
||||
```
|
||||
|
||||
2. **添加显式刷新**:
|
||||
```python
|
||||
# 第496行后添加
|
||||
if writer:
|
||||
writer.flush() # 确保数据写入磁盘
|
||||
```
|
||||
|
||||
3. **添加初始化验证**:
|
||||
```python
|
||||
# 第271行后添加
|
||||
if writer:
|
||||
# 测试写入
|
||||
writer.add_scalar('Test/Initialization', 0.0, 0)
|
||||
writer.flush()
|
||||
print(f"✓ TensorBoard logging enabled. Log dir: {args.log_dir}")
|
||||
# 检查文件是否创建
|
||||
import glob
|
||||
event_files = glob.glob(os.path.join(args.log_dir, "events.out.tfevents.*"))
|
||||
if event_files:
|
||||
print(f"✓ TensorBoard event file created: {event_files[0]}")
|
||||
else:
|
||||
print("⚠ WARNING: TensorBoard not available. Logging disabled.")
|
||||
```
|
||||
|
||||
4. **在训练结束时确保关闭**:
|
||||
```python
|
||||
# 第505行后添加
|
||||
if writer:
|
||||
writer.flush() # 最后一次刷新
|
||||
writer.close()
|
||||
print(f"TensorBoard logs saved to {args.log_dir}")
|
||||
```
|
||||
|
||||
### 5.5 验证修复
|
||||
|
||||
**重新训练测试**:
|
||||
```bash
|
||||
python train_magail.py \
|
||||
--expert_data_dir data/training_data \
|
||||
--data_dir data/exp_filtered \
|
||||
--batch_size 1024 \
|
||||
--max_episodes 10 \
|
||||
--log_dir runs/test_tensorboard
|
||||
```
|
||||
|
||||
**检查输出**:
|
||||
```bash
|
||||
# 应该看到事件文件
|
||||
ls runs/test_tensorboard/events.out.tfevents.*
|
||||
|
||||
# 启动 TensorBoard
|
||||
tensorboard --logdir runs/test_tensorboard
|
||||
```
|
||||
|
||||
**对于 magail_0112 训练**:
|
||||
由于该训练已经完成且没有日志文件,建议:
|
||||
1. 检查训练时的控制台输出,确认是否有 "TensorBoard not installed" 消息
|
||||
2. 如果确实没有 TensorBoard,可以重新运行少量 episode 来验证修复
|
||||
3. 或者查看是否有其他日志文件(如 `training.log`)
|
||||
|
||||
---
|
||||
|
||||
## 附录: 关键文件清单
|
||||
|
||||
### 核心训练文件
|
||||
- `train_magail.py`: 主训练脚本
|
||||
- `dataset/magail_dataset.py`: 专家数据集加载
|
||||
- `Env/expert_replay_env.py`: 专家回放环境
|
||||
- `Env/scenario_env.py`: 多智能体场景环境
|
||||
- `Env/inverse_dynamics.py`: 逆动力学计算
|
||||
|
||||
### 数据生成文件
|
||||
- `scripts/generate_expert_data.py`: 专家数据生成
|
||||
- `scripts/visualize_replay.py`: 数据可视化
|
||||
- `scripts/analyze_expert_data.py`: 数据分析
|
||||
|
||||
### 配置文件
|
||||
- `README.md`: 项目说明
|
||||
- `TRAINING_ARCHITECTURE.md`: 本文档
|
||||
|
||||
---
|
||||
|
||||
## 总结
|
||||
|
||||
本项目的 MAGAIL 训练方案通过以下方式处理多智能体可变长度问题:
|
||||
|
||||
1. **数据层面**: 将轨迹展平为独立样本,统一处理
|
||||
2. **环境层面**: 动态处理不同数量的智能体,合并经验
|
||||
3. **网络层面**: 固定输入维度(45维),通过零填充处理邻居不足的情况
|
||||
|
||||
训练流程遵循标准的 GAIL 框架,使用 PPO 作为策略优化算法。当前主要问题集中在环境稳定性和日志记录方面,需要进一步优化。
|
||||
18
docs/examples/hbbc_latent_example.json
Normal file
18
docs/examples/hbbc_latent_example.json
Normal file
@@ -0,0 +1,18 @@
|
||||
{
|
||||
"global": {
|
||||
"latent_eps": [0.35, -0.12, 0.28, 0.46, -0.22, 0.18],
|
||||
"latent_c": [0, 1,1, 0]
|
||||
},
|
||||
"object_id": {
|
||||
"12345": {
|
||||
"latent_eps": [0.2, -0.1, 0.3, 0.4, -0.2, 0.1],
|
||||
"latent_c": [0, 1, 0, 0]
|
||||
}
|
||||
},
|
||||
"agent_id": {
|
||||
"controlled_abcde": {
|
||||
"latent_eps": [0.5, 0.1, -0.1, 0.2, -0.3, 0.4],
|
||||
"latent_c": [1, 0, 0, 0]
|
||||
}
|
||||
}
|
||||
}
|
||||
BIN
expert_trajectories_full.pkl
Normal file
BIN
expert_trajectories_full.pkl
Normal file
Binary file not shown.
BIN
expert_trajectories_full_obs.pkl
Normal file
BIN
expert_trajectories_full_obs.pkl
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
105
scripts/README.md
Normal file
105
scripts/README.md
Normal file
@@ -0,0 +1,105 @@
|
||||
# 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 | 见下方 |
|
||||
|
||||
**多智能体**(输出 `expert_data_{start_index}_{num_scenarios}.pkl`):
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||
```
|
||||
|
||||
**单智能体**(仅采集 ego 车轨迹,输出 `expert_data_ego_{start_index}_{num_scenarios}.pkl`,用于单智能体 BC):
|
||||
```bash
|
||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||
```
|
||||
|
||||
**常用参数**:`--data_dir`(默认 `data/exp_filtered`)、`--output_dir`(默认 `data/training_data`)、`--start_index`、`--num_scenarios`、`--ego_only`(仅保存 default_agent 轨迹,输出使用 `expert_data_ego_*.pkl` 前缀)。
|
||||
|
||||
---
|
||||
|
||||
### 可视化(统一入口)
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [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 训练策略):与专家数据生成/回放一致——同一套车道+静态筛选、且会生成背景车(bg_*),使观测分布与训练集一致,便于在训练集上公平演示。
|
||||
```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
|
||||
```
|
||||
- **policy + 仅自车策略、其他车回放**(BC 单智能体模型):加 `--ego_only`,自车由策略控制,其余车辆按专家轨迹回放。
|
||||
```bash
|
||||
python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1 --ego_only
|
||||
```
|
||||
|
||||
- **policy + HBBC 动态背景车**(仅动态背景车启用,静态背景车保持原样):
|
||||
```bash
|
||||
python scripts/visualize.py policy \
|
||||
--policy_type bc \
|
||||
--model_path models/bc/policy_best.pt \
|
||||
--data_dir data/exp_filtered \
|
||||
--num_scenarios 1 \
|
||||
--ego_only \
|
||||
--enable_hbbc_background \
|
||||
--hbbc_model_path models/hbbc/hbbc.pt \
|
||||
--hbbc_inference_device cpu \
|
||||
--hbbc_latent_mode per_vehicle_fixed \
|
||||
--hbbc_latent_json_path docs/examples/hbbc_latent_example.json
|
||||
```
|
||||
|
||||
- **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)、`--ego_only`(仅 BC:自车用策略,其他车专家回放)、`--enable_hbbc_background`、`--hbbc_model_path`、`--hbbc_inference_device`、`--hbbc_latent_mode`、`--hbbc_latent_json_path`。
|
||||
|
||||
---
|
||||
|
||||
### 数据分析与检查
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [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`(多智能体 `expert_data_*.pkl`,单智能体 `expert_data_ego_*.pkl`)
|
||||
2. **BC 训练**:根目录 `train_bc.py` → 模型保存到 `models/bc/`,日志到 `logs/bc/`。单智能体模式加 `--single_agent` 并指定 ego-only 的 pkl。
|
||||
3. **MAGAIL 训练**:根目录 `train_magail.py` → 模型保存到 `models/magail/`,日志到 `logs/magail/`
|
||||
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
||||
0
scripts/__init__.py
Normal file
0
scripts/__init__.py
Normal file
256
scripts/analyze_expert_data.py
Normal file
256
scripts/analyze_expert_data.py
Normal file
@@ -0,0 +1,256 @@
|
||||
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)
|
||||
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
from collections import defaultdict
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
import pickle
|
||||
import os
|
||||
|
||||
class DummyPolicy:
|
||||
"""占位策略"""
|
||||
def act(self, *args, **kwargs):
|
||||
return np.array([0.0, 0.0])
|
||||
|
||||
class ExpertDataAnalyzer:
|
||||
def __init__(self, data_directory):
|
||||
self.data_directory = data_directory
|
||||
self.env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": data_directory,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
},
|
||||
agent2policy=DummyPolicy() # 添加必需参数
|
||||
)
|
||||
|
||||
self.statistics = {
|
||||
"num_scenarios": 0,
|
||||
"num_trajectories": 0,
|
||||
"trajectory_lengths": [],
|
||||
"velocities": [],
|
||||
"speeds": [], # 速度大小
|
||||
"accelerations": [],
|
||||
"heading_changes": [],
|
||||
"inter_vehicle_distances": [],
|
||||
"num_vehicles_per_scenario": [],
|
||||
"static_vehicles": 0, # 统计静止车辆
|
||||
}
|
||||
|
||||
def analyze_all_scenarios(self, num_scenarios=None):
|
||||
"""遍历所有场景并收集统计信息"""
|
||||
scenario_count = 0
|
||||
|
||||
while True:
|
||||
try:
|
||||
obs = self.env.reset()
|
||||
|
||||
if not hasattr(self.env, 'expert_trajectories'):
|
||||
print("⚠️ 环境缺少expert_trajectories属性")
|
||||
break
|
||||
|
||||
expert_trajs = self.env.expert_trajectories
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
continue
|
||||
|
||||
scenario_count += 1
|
||||
self.statistics["num_scenarios"] += 1
|
||||
self.statistics["num_vehicles_per_scenario"].append(len(expert_trajs))
|
||||
|
||||
# 分析每条轨迹
|
||||
for obj_id, traj in expert_trajs.items():
|
||||
self.analyze_single_trajectory(traj)
|
||||
|
||||
# 分析车辆间交互
|
||||
self.analyze_vehicle_interactions(expert_trajs)
|
||||
|
||||
print(f"已分析场景 {scenario_count}/{num_scenarios}, 车辆数: {len(expert_trajs)}")
|
||||
|
||||
if num_scenarios and scenario_count >= num_scenarios:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
print(f"场景 {scenario_count} 处理失败: {e}")
|
||||
break
|
||||
|
||||
self.env.close()
|
||||
|
||||
def analyze_single_trajectory(self, traj):
|
||||
"""分析单条轨迹"""
|
||||
self.statistics["num_trajectories"] += 1
|
||||
|
||||
length = traj["length"]
|
||||
self.statistics["trajectory_lengths"].append(length)
|
||||
|
||||
# 速度分析
|
||||
velocities = traj["velocities"]
|
||||
speeds = np.linalg.norm(velocities, axis=1)
|
||||
self.statistics["velocities"].extend(velocities.tolist())
|
||||
self.statistics["speeds"].extend(speeds.tolist())
|
||||
|
||||
# 检查是否为静止车辆
|
||||
if np.max(speeds) < 0.5: # 最大速度小于0.5m/s视为静止
|
||||
self.statistics["static_vehicles"] += 1
|
||||
|
||||
# 加速度分析
|
||||
if length > 1:
|
||||
accelerations = np.diff(speeds) * 10 # 10Hz数据
|
||||
self.statistics["accelerations"].extend(accelerations.tolist())
|
||||
|
||||
# 航向角变化
|
||||
headings = traj["headings"]
|
||||
if length > 1:
|
||||
heading_changes = np.diff(headings)
|
||||
heading_changes = np.arctan2(np.sin(heading_changes), np.cos(heading_changes))
|
||||
self.statistics["heading_changes"].extend(heading_changes.tolist())
|
||||
|
||||
def analyze_vehicle_interactions(self, expert_trajs):
|
||||
"""分析车辆间的距离"""
|
||||
if len(expert_trajs) < 2:
|
||||
return
|
||||
|
||||
traj_list = list(expert_trajs.values())
|
||||
|
||||
for i in range(len(traj_list)):
|
||||
for j in range(i+1, len(traj_list)):
|
||||
traj_i = traj_list[i]
|
||||
traj_j = traj_list[j]
|
||||
|
||||
start_time = max(traj_i["start_timestep"], traj_j["start_timestep"])
|
||||
end_time = min(traj_i["end_timestep"], traj_j["end_timestep"])
|
||||
|
||||
if start_time >= end_time:
|
||||
continue
|
||||
|
||||
idx_i_start = start_time - traj_i["start_timestep"]
|
||||
idx_i_end = end_time - traj_i["start_timestep"]
|
||||
idx_j_start = start_time - traj_j["start_timestep"]
|
||||
idx_j_end = end_time - traj_j["start_timestep"]
|
||||
|
||||
pos_i = traj_i["positions"][idx_i_start:idx_i_end, :2]
|
||||
pos_j = traj_j["positions"][idx_j_start:idx_j_end, :2]
|
||||
|
||||
distances = np.linalg.norm(pos_i - pos_j, axis=1)
|
||||
self.statistics["inter_vehicle_distances"].extend(distances.tolist())
|
||||
|
||||
def generate_report(self, save_dir="./analysis_results"):
|
||||
"""生成统计报告"""
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
stats = self.statistics
|
||||
|
||||
print("\n" + "="*60)
|
||||
print("专家数据集统计报告")
|
||||
print("="*60)
|
||||
print(f"总场景数: {stats['num_scenarios']}")
|
||||
print(f"总轨迹数: {stats['num_trajectories']}")
|
||||
print(f"静止车辆数: {stats['static_vehicles']} ({stats['static_vehicles']/stats['num_trajectories']*100:.1f}%)")
|
||||
print(f"平均每场景车辆数: {np.mean(stats['num_vehicles_per_scenario']):.2f} ± {np.std(stats['num_vehicles_per_scenario']):.2f}")
|
||||
|
||||
print(f"\n轨迹长度统计 (帧数 @ 10Hz):")
|
||||
print(f" 平均: {np.mean(stats['trajectory_lengths']):.2f} 帧 ({np.mean(stats['trajectory_lengths'])*0.1:.2f}秒)")
|
||||
print(f" 中位数: {np.median(stats['trajectory_lengths']):.2f} 帧")
|
||||
print(f" 最小/最大: {np.min(stats['trajectory_lengths'])} / {np.max(stats['trajectory_lengths'])} 帧")
|
||||
|
||||
print(f"\n速度统计 (m/s):")
|
||||
speeds = np.array(stats['speeds'])
|
||||
print(f" 平均: {np.mean(speeds):.2f} ± {np.std(speeds):.2f}")
|
||||
print(f" 中位数: {np.median(speeds):.2f}")
|
||||
print(f" 最小/最大: {np.min(speeds):.2f} / {np.max(speeds):.2f}")
|
||||
print(f" 静止帧(<0.5m/s): {np.sum(speeds < 0.5)} ({np.sum(speeds < 0.5)/len(speeds)*100:.1f}%)")
|
||||
|
||||
print(f"\n加速度统计 (m/s²):")
|
||||
accs = np.array(stats['accelerations'])
|
||||
print(f" 平均: {np.mean(accs):.4f} ± {np.std(accs):.2f}")
|
||||
print(f" 最小/最大: {np.min(accs):.2f} / {np.max(accs):.2f}")
|
||||
|
||||
if len(stats['inter_vehicle_distances']) > 0:
|
||||
dists = np.array(stats['inter_vehicle_distances'])
|
||||
print(f"\n车辆间距离统计 (m):")
|
||||
print(f" 平均: {np.mean(dists):.2f} ± {np.std(dists):.2f}")
|
||||
print(f" 最小: {np.min(dists):.2f}")
|
||||
print(f" 近距离交互(<5m): {np.sum(dists < 5.0)} ({np.sum(dists < 5.0)/len(dists)*100:.2f}%)")
|
||||
|
||||
# 保存数据
|
||||
with open(os.path.join(save_dir, "statistics.pkl"), "wb") as f:
|
||||
pickle.dump(stats, f)
|
||||
|
||||
# 绘制可视化
|
||||
self.plot_distributions(save_dir)
|
||||
|
||||
print(f"\n✓ 报告已保存到: {save_dir}")
|
||||
|
||||
def plot_distributions(self, save_dir):
|
||||
"""绘制分布图"""
|
||||
stats = self.statistics
|
||||
|
||||
fig, axes = plt.subplots(2, 3, figsize=(15, 10))
|
||||
|
||||
# 1. 轨迹长度分布
|
||||
axes[0, 0].hist(stats['trajectory_lengths'], bins=50, edgecolor='black')
|
||||
axes[0, 0].set_xlabel('Trajectory Length (frames @ 10Hz)')
|
||||
axes[0, 0].set_ylabel('Frequency')
|
||||
axes[0, 0].set_title('Trajectory Length Distribution')
|
||||
axes[0, 0].axvline(np.mean(stats['trajectory_lengths']), color='red',
|
||||
linestyle='--', label=f'Mean: {np.mean(stats["trajectory_lengths"]):.1f}')
|
||||
axes[0, 0].legend()
|
||||
|
||||
# 2. 速度分布
|
||||
axes[0, 1].hist(stats['speeds'], bins=50, edgecolor='black')
|
||||
axes[0, 1].set_xlabel('Speed (m/s)')
|
||||
axes[0, 1].set_ylabel('Frequency')
|
||||
axes[0, 1].set_title('Speed Distribution')
|
||||
axes[0, 1].axvline(np.mean(stats['speeds']), color='red',
|
||||
linestyle='--', label=f'Mean: {np.mean(stats["speeds"]):.2f}')
|
||||
axes[0, 1].legend()
|
||||
|
||||
# 3. 加速度分布
|
||||
axes[0, 2].hist(stats['accelerations'], bins=50, edgecolor='black')
|
||||
axes[0, 2].set_xlabel('Acceleration (m/s²)')
|
||||
axes[0, 2].set_ylabel('Frequency')
|
||||
axes[0, 2].set_title('Acceleration Distribution')
|
||||
|
||||
# 4. 每场景车辆数
|
||||
axes[1, 0].hist(stats['num_vehicles_per_scenario'], bins=30, edgecolor='black')
|
||||
axes[1, 0].set_xlabel('Vehicles per Scenario')
|
||||
axes[1, 0].set_ylabel('Frequency')
|
||||
axes[1, 0].set_title('Vehicles per Scenario')
|
||||
|
||||
# 5. 航向角变化
|
||||
axes[1, 1].hist(stats['heading_changes'], bins=50, edgecolor='black')
|
||||
axes[1, 1].set_xlabel('Heading Change (rad)')
|
||||
axes[1, 1].set_ylabel('Frequency')
|
||||
axes[1, 1].set_title('Heading Change Distribution')
|
||||
|
||||
# 6. 车辆间距离
|
||||
if len(stats['inter_vehicle_distances']) > 0:
|
||||
axes[1, 2].hist(stats['inter_vehicle_distances'], bins=50,
|
||||
range=(0, 50), edgecolor='black')
|
||||
axes[1, 2].set_xlabel('Inter-vehicle Distance (m)')
|
||||
axes[1, 2].set_ylabel('Frequency')
|
||||
axes[1, 2].set_title('Distance Distribution')
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig(os.path.join(save_dir, "distributions.png"), dpi=300)
|
||||
print(f" ✓ 分布图已保存")
|
||||
|
||||
if __name__ == "__main__":
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/data"
|
||||
data_dir = AssetLoader.file_path(WAYMO_DATA_DIR, "exp_filtered", unix_style=False)
|
||||
|
||||
print("开始分析专家数据...")
|
||||
analyzer = ExpertDataAnalyzer(data_dir)
|
||||
analyzer.analyze_all_scenarios(num_scenarios=100) # 分析100个场景
|
||||
analyzer.generate_report()
|
||||
47
scripts/check_database_info.py
Normal file
47
scripts/check_database_info.py
Normal file
@@ -0,0 +1,47 @@
|
||||
import pickle
|
||||
import os
|
||||
|
||||
# 检查过滤后的数据库
|
||||
filtered_db = "/home/huangfukk/mdsn/exp_filtered"
|
||||
|
||||
print("="*60)
|
||||
print("过滤后数据库信息")
|
||||
print("="*60)
|
||||
|
||||
# 读取summary
|
||||
summary_path = os.path.join(filtered_db, "dataset_summary.pkl")
|
||||
with open(summary_path, 'rb') as f:
|
||||
summary = pickle.load(f)
|
||||
|
||||
print(f"\n总场景数: {len(summary)}")
|
||||
print(f"场景ID列表(前10个): {list(summary.keys())[:10]}")
|
||||
|
||||
# 读取mapping
|
||||
mapping_path = os.path.join(filtered_db, "dataset_mapping.pkl")
|
||||
with open(mapping_path, 'rb') as f:
|
||||
mapping = pickle.load(f)
|
||||
|
||||
print(f"\n映射关系数量: {len(mapping)}")
|
||||
|
||||
# 检查第一个场景的详细信息
|
||||
first_scenario_id = list(summary.keys())[0]
|
||||
first_scenario_info = summary[first_scenario_id]
|
||||
print(f"\n第一个场景详细信息:")
|
||||
print(f" 场景ID: {first_scenario_id}")
|
||||
print(f" 元数据: {first_scenario_info}")
|
||||
|
||||
# 检查映射的文件路径
|
||||
first_scenario_path = mapping[first_scenario_id]
|
||||
print(f" 场景文件路径(相对): {first_scenario_path}")
|
||||
|
||||
# 检查文件是否存在
|
||||
abs_path = os.path.join(filtered_db, first_scenario_path)
|
||||
print(f" 场景文件路径(绝对): {abs_path}")
|
||||
print(f" 文件存在: {os.path.exists(abs_path)}")
|
||||
|
||||
# 统计源数据库的场景文件
|
||||
converted_db = "/home/huangfukk/mdsn/exp_converted"
|
||||
converted_files = [f for f in os.listdir(converted_db) if f.endswith('.pkl') and f.startswith('sd_')]
|
||||
print(f"\n源数据库 exp_converted:")
|
||||
print(f" 场景文件数量: {len(converted_files)}")
|
||||
print(f" 示例文件: {converted_files[:5]}")
|
||||
177
scripts/check_track_fields.py
Normal file
177
scripts/check_track_fields.py
Normal file
@@ -0,0 +1,177 @@
|
||||
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
|
||||
|
||||
class DummyPolicy:
|
||||
"""
|
||||
占位策略,用于数据检查时初始化环境
|
||||
不需要实际执行动作,只是为了满足环境初始化要求
|
||||
"""
|
||||
def act(self, *args, **kwargs):
|
||||
# 返回零动作 [throttle, steering]
|
||||
return np.array([0.0, 0.0])
|
||||
|
||||
def check_available_fields():
|
||||
"""
|
||||
检查Waymo转MetaDrive数据中实际可用的字段
|
||||
"""
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/mdsn"
|
||||
data_dir = AssetLoader.file_path(WAYMO_DATA_DIR, "exp_filtered", unix_style=False)
|
||||
|
||||
# 创建占位策略
|
||||
dummy_policy = DummyPolicy()
|
||||
|
||||
# 初始化环境,传入必需的agent2policy参数
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
},
|
||||
agent2policy=dummy_policy # 添加这个必需参数
|
||||
)
|
||||
|
||||
print("✓ 环境初始化成功")
|
||||
|
||||
# 重置环境以加载数据
|
||||
print("正在加载场景数据...")
|
||||
env.reset()
|
||||
|
||||
# 检查是否有expert_trajectories属性
|
||||
if hasattr(env, 'expert_trajectories'):
|
||||
print(f"✓ expert_trajectories属性存在,包含 {len(env.expert_trajectories)} 条轨迹")
|
||||
else:
|
||||
print("⚠️ expert_trajectories属性不存在,请先修改scenario_env.py添加轨迹存储功能")
|
||||
|
||||
# 获取一个track样本
|
||||
sample_track = None
|
||||
for scenario_id, track in env.engine.traffic_manager.current_traffic_data.items():
|
||||
if track["type"] == "VEHICLE":
|
||||
sample_track = track
|
||||
print(f"\n找到样本车辆: scenario_id = {scenario_id}")
|
||||
break
|
||||
|
||||
if sample_track is None:
|
||||
print("未找到车辆轨迹数据")
|
||||
env.close()
|
||||
return
|
||||
|
||||
print("="*60)
|
||||
print("Track数据结构分析")
|
||||
print("="*60)
|
||||
|
||||
# 1. 顶层字段
|
||||
print("\n1. Track顶层字段:")
|
||||
for key in sample_track.keys():
|
||||
print(f" - {key}: {type(sample_track[key])}")
|
||||
|
||||
# 2. metadata字段
|
||||
print("\n2. track['metadata']字段:")
|
||||
if "metadata" in sample_track:
|
||||
for key, value in sample_track["metadata"].items():
|
||||
if isinstance(value, (str, int, float, bool)):
|
||||
print(f" - {key}: {type(value).__name__} = {value}")
|
||||
else:
|
||||
print(f" - {key}: {type(value).__name__}")
|
||||
|
||||
# 3. state字段
|
||||
print("\n3. track['state']字段:")
|
||||
if "state" in sample_track:
|
||||
for key, value in sample_track["state"].items():
|
||||
if isinstance(value, np.ndarray):
|
||||
print(f" - {key}: shape={value.shape}, dtype={value.dtype}")
|
||||
# 打印第一个有效值
|
||||
if "valid" in sample_track["state"]:
|
||||
valid_idx = np.argmax(sample_track["state"]["valid"])
|
||||
if valid_idx >= 0 and valid_idx < len(value):
|
||||
print(f" 示例值 (index {valid_idx}): {value[valid_idx]}")
|
||||
else:
|
||||
print(f" - {key}: {type(value)} = {value}")
|
||||
|
||||
print("\n" + "="*60)
|
||||
print("建议存储的字段:")
|
||||
print("="*60)
|
||||
|
||||
# 检查必需字段
|
||||
required_fields = ["position", "heading", "velocity", "valid"]
|
||||
print("\n必需字段:")
|
||||
all_required_exist = True
|
||||
for field in required_fields:
|
||||
if "state" in sample_track and field in sample_track["state"]:
|
||||
print(f" ✓ {field} (存在)")
|
||||
else:
|
||||
print(f" ✗ {field} (缺失)")
|
||||
all_required_exist = False
|
||||
|
||||
# 检查可选字段
|
||||
optional_fields = ["length", "width", "height", "bbox"]
|
||||
print("\n可选字段:")
|
||||
available_optional = []
|
||||
for field in optional_fields:
|
||||
if "state" in sample_track and field in sample_track["state"]:
|
||||
print(f" + {field} (在state中)")
|
||||
available_optional.append(field)
|
||||
elif "metadata" in sample_track and field in sample_track["metadata"]:
|
||||
print(f" + {field} (在metadata中)")
|
||||
available_optional.append(field)
|
||||
else:
|
||||
print(f" - {field} (不存在)")
|
||||
|
||||
print("\n" + "="*60)
|
||||
print("推荐的trajectory_data结构:")
|
||||
print("="*60)
|
||||
|
||||
if all_required_exist:
|
||||
print("""
|
||||
trajectory_data = {
|
||||
"object_id": object_id,
|
||||
"scenario_id": scenario_id,
|
||||
"valid_mask": valid[first_show:last_show+1].copy(),
|
||||
"positions": track["state"]["position"][first_show:last_show+1].copy(),
|
||||
"headings": track["state"]["heading"][first_show:last_show+1].copy(),
|
||||
"velocities": track["state"]["velocity"][first_show:last_show+1].copy(),
|
||||
"timesteps": np.arange(first_show, last_show+1),
|
||||
"start_timestep": first_show,
|
||||
"end_timestep": last_show,
|
||||
"length": last_show - first_show + 1
|
||||
}
|
||||
""")
|
||||
|
||||
if available_optional:
|
||||
print("如果需要车辆尺寸,可选添加:")
|
||||
for field in available_optional:
|
||||
if field in ["length", "width", "height"]:
|
||||
print(f' trajectory_data["vehicle_{field}"] = track["state" or "metadata"]["{field}"][first_show]')
|
||||
else:
|
||||
print("⚠️ 缺少必需字段,请检查数据转换流程")
|
||||
|
||||
# 如果有expert_trajectories,展示一个样本
|
||||
if hasattr(env, 'expert_trajectories') and len(env.expert_trajectories) > 0:
|
||||
print("\n" + "="*60)
|
||||
print("expert_trajectories样本:")
|
||||
print("="*60)
|
||||
sample_traj = list(env.expert_trajectories.values())[0]
|
||||
for key, value in sample_traj.items():
|
||||
if isinstance(value, np.ndarray):
|
||||
print(f" {key}: shape={value.shape}, dtype={value.dtype}")
|
||||
else:
|
||||
print(f" {key}: {type(value).__name__} = {value}")
|
||||
|
||||
env.close()
|
||||
print("\n✓ 分析完成")
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_available_fields()
|
||||
169
scripts/generate_expert_data.py
Normal file
169
scripts/generate_expert_data.py
Normal file
@@ -0,0 +1,169 @@
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import pickle
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
# 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 metadrive.engine.asset_loader import AssetLoader
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
def generate_data(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")
|
||||
|
||||
# MetaDrive's ScenarioDataManager asserts if config["num_scenarios"] > available scenarios in data_directory.
|
||||
# So we always set it to -1 (load all available) and clamp the loop range by reading dataset summary.
|
||||
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, # Set high to catch all vehicles in scenario
|
||||
"horizon": 1000,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False, # Important: we replay, not react
|
||||
"start_scenario_index": args.start_index,
|
||||
# Load all scenarios available in the directory to avoid assertion failure.
|
||||
# We will still only iterate `num_to_run` scenarios below.
|
||||
"num_scenarios": -1,
|
||||
"log_level": 50 # ERROR to reduce noise
|
||||
}
|
||||
|
||||
expert_trajectories = []
|
||||
|
||||
try:
|
||||
# Loop through scenarios
|
||||
for i in tqdm(range(args.start_index, args.start_index + num_to_run), desc="Scenarios"):
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
try:
|
||||
obs_dict = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
try:
|
||||
env.close()
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
|
||||
# Storage for current episode
|
||||
# dict of lists: {agent_id: {'obs': [], 'acts': []}}
|
||||
episode_data = {}
|
||||
|
||||
# Map agent_id to original ID if possible, but agent_id is unique enough
|
||||
|
||||
for step in range(env.config["horizon"]):
|
||||
# Step with dummy actions
|
||||
obs, rewards, dones, infos = env.step(None)
|
||||
|
||||
# 'obs' is next observation (t+1)
|
||||
# 'infos' contains 'expert_action' which took (t -> t+1)
|
||||
# Wait, usually (obs_t, act_t) -> obs_{t+1}
|
||||
# expert_replay_env.step():
|
||||
# calc action (t -> t+1)
|
||||
# move agents to t+1
|
||||
# return obs_{t+1}
|
||||
# So we have obs_dict (from reset or prev step) which is at 't'
|
||||
# And we have 'infos' which has action at 't'.
|
||||
|
||||
current_agents = list(obs_dict.keys())
|
||||
|
||||
for agent_id in current_agents:
|
||||
if agent_id not in episode_data:
|
||||
episode_data[agent_id] = {'obs': [], 'acts': []}
|
||||
|
||||
# Check if we have action for this agent
|
||||
if agent_id in infos and 'expert_action' in infos[agent_id]:
|
||||
action = infos[agent_id]['expert_action']
|
||||
observation = obs_dict[agent_id]
|
||||
|
||||
episode_data[agent_id]['obs'].append(observation)
|
||||
episode_data[agent_id]['acts'].append(action)
|
||||
|
||||
# Update obs_dict for next step
|
||||
obs_dict = obs
|
||||
|
||||
if dones["__all__"]:
|
||||
break
|
||||
|
||||
# Post-process episode data
|
||||
for agent_id, data in episode_data.items():
|
||||
if args.ego_only and agent_id != "default_agent":
|
||||
continue
|
||||
if len(data['obs']) > 10: # Minimum length filter
|
||||
expert_trajectories.append({
|
||||
'obs': np.array(data['obs']),
|
||||
'acts': np.array(data['acts']),
|
||||
'agent_id': agent_id,
|
||||
'scenario_id': i
|
||||
})
|
||||
env.close()
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
# env is closed per-scenario above (more robust for MetaDrive object lifecycle)
|
||||
pass
|
||||
|
||||
# Save data
|
||||
if args.ego_only:
|
||||
output_file = os.path.join(args.output_dir, f"expert_data_ego_{args.start_index}_{args.num_scenarios}.pkl")
|
||||
else:
|
||||
output_file = os.path.join(args.output_dir, f"expert_data_{args.start_index}_{args.num_scenarios}.pkl")
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
if args.ego_only:
|
||||
print("Ego-only mode: saved trajectories are SDC (default_agent) only.")
|
||||
print(f"Saving {len(expert_trajectories)} trajectories to {output_file}")
|
||||
with open(output_file, 'wb') as f:
|
||||
pickle.dump(expert_trajectories, f)
|
||||
|
||||
# Verification stats
|
||||
if len(expert_trajectories) > 0:
|
||||
all_acts = np.concatenate([t['acts'] for t in expert_trajectories])
|
||||
print("Action Stats:")
|
||||
print(f" Steering: min={all_acts[:,0].min():.3f}, max={all_acts[:,0].max():.3f}, mean={all_acts[:,0].mean():.3f}")
|
||||
print(f" Accel: min={all_acts[:,1].min():.3f}, max={all_acts[:,1].max():.3f}, mean={all_acts[:,1].mean():.3f}")
|
||||
|
||||
# Clipping ratio diagnostics (actions are normalized to [-1, 1])
|
||||
# If this ratio is high, it usually indicates max_acc/max_steering too small or noisy finite-difference.
|
||||
eps = 1e-6
|
||||
steer = all_acts[:, 0]
|
||||
accel = all_acts[:, 1]
|
||||
steer_clipped = np.isclose(np.abs(steer), 1.0, atol=eps)
|
||||
accel_clipped = np.isclose(np.abs(accel), 1.0, atol=eps)
|
||||
print("Clipping Stats:")
|
||||
print(
|
||||
f" Steering clipped (|a|==1): {steer_clipped.mean()*100:.2f}% "
|
||||
f"({steer_clipped.sum()}/{len(steer_clipped)})"
|
||||
)
|
||||
print(
|
||||
f" Accel clipped (|a|==1): {accel_clipped.mean()*100:.2f}% "
|
||||
f"({accel_clipped.sum()}/{len(accel_clipped)})"
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
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)
|
||||
parser.add_argument("--ego_only", action="store_true", help="Only collect and save ego (default_agent) trajectories; output uses expert_data_ego_*.pkl prefix")
|
||||
args = parser.parse_args()
|
||||
generate_data(args)
|
||||
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())
|
||||
432
scripts/visualize.py
Normal file
432
scripts/visualize.py
Normal file
@@ -0,0 +1,432 @@
|
||||
"""
|
||||
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)
|
||||
|
||||
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + num_to_run):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
# Each scenario uses a fresh env (start_scenario_index=i, num_scenarios=1) so the second
|
||||
# scenario and beyond are fully cleaned and loaded like a single-scenario run.
|
||||
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": i,
|
||||
"num_scenarios": 1,
|
||||
"log_level": 40,
|
||||
}
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
try:
|
||||
obs = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
env.close()
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents (current): {len(env.controlled_agents)}, total in scenario: {env.num_controlled_in_scenario}")
|
||||
|
||||
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
|
||||
env.close()
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted by user")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
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 Env.bc_ego_replay_env import BCEgoReplayEnv
|
||||
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"
|
||||
ego_only = getattr(args, "ego_only", False)
|
||||
if ego_only and policy_type != "bc":
|
||||
print("[WARN] --ego_only is supported for BC policy only; MAGAIL will run in multi-agent mode.")
|
||||
|
||||
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 if ego_only else 3,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True,
|
||||
"sequential_seed": True,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": args.num_scenarios,
|
||||
"log_level": 40,
|
||||
"enable_hbbc_background": bool(getattr(args, "enable_hbbc_background", False)),
|
||||
"hbbc_model_path": getattr(args, "hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||
"hbbc_inference_device": getattr(args, "hbbc_inference_device", "cpu"),
|
||||
"hbbc_latent_mode": getattr(args, "hbbc_latent_mode", "per_vehicle_fixed"),
|
||||
"hbbc_latent_json_path": getattr(args, "hbbc_latent_json_path", None),
|
||||
}
|
||||
|
||||
if ego_only and policy_type == "bc":
|
||||
print("Initializing BCEgoReplayEnv (ego-only: policy on self, others replayed)...")
|
||||
else:
|
||||
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
||||
|
||||
try:
|
||||
env = BCEgoReplayEnv(config=env_config) if (ego_only and policy_type == "bc") else 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 = BCEgoReplayEnv(config=env_config) if (ego_only and policy_type == "bc") else 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)
|
||||
try:
|
||||
state = torch.load(model_path, map_location=device, weights_only=True)
|
||||
except TypeError:
|
||||
state = torch.load(model_path, map_location=device)
|
||||
policy.load_state_dict(state)
|
||||
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
|
||||
|
||||
n_total = getattr(env, "num_controlled_in_scenario", len(obs_dict))
|
||||
mode_note = " (ego only, others replayed)" if (ego_only and policy_type == "bc") else ""
|
||||
if ego_only and policy_type == "bc" and bool(env_config.get("enable_hbbc_background", False)):
|
||||
mode_note = " (ego only, dynamic background via HBBC)"
|
||||
print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {n_total}{mode_note}")
|
||||
if ego_only and policy_type == "bc" and len(obs_dict) == 1:
|
||||
if bool(env_config.get("enable_hbbc_background", False)):
|
||||
print(" [Ego control: policy injected — dynamic background vehicles use HBBC; static background stays static.]")
|
||||
else:
|
||||
print(" [Ego control: policy injected — ego uses model output each step; other vehicles expert replay.]")
|
||||
if len(obs_dict) == 0:
|
||||
print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.")
|
||||
continue
|
||||
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")
|
||||
pp.add_argument("--ego_only", action="store_true", help="BC only: inject policy into ego only; other vehicles use expert replay")
|
||||
pp.add_argument("--enable_hbbc_background", action="store_true", help="Enable HBBC policy for dynamic background vehicles")
|
||||
pp.add_argument("--hbbc_model_path", type=str, default="models/hbbc/hbbc.pt")
|
||||
pp.add_argument("--hbbc_inference_device", type=str, default="cpu")
|
||||
pp.add_argument("--hbbc_latent_mode", type=str, default="per_vehicle_fixed", choices=["per_vehicle_fixed", "per_episode_reset"])
|
||||
pp.add_argument("--hbbc_latent_json_path", type=str, default=None, help="Optional JSON for per-vehicle latent override")
|
||||
|
||||
# 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()
|
||||
212
train_bc.py
Normal file
212
train_bc.py
Normal file
@@ -0,0 +1,212 @@
|
||||
"""
|
||||
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 Env.bc_ego_replay_env import BCEgoReplayEnv
|
||||
from dataset.loader import load_expert_pkl, get_expert_scenario_ids
|
||||
|
||||
|
||||
def evaluate_policy(policy, args, device):
|
||||
"""在 BCScenarioEnv(多智能体)或 BCEgoReplayEnv(单智能体)中评估策略。
|
||||
仅使用专家数据中出现过的 scenario_id。单智能体模式下仅 ego 受策略控制,其他车专家回放。"""
|
||||
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, 0.0, 0.0
|
||||
|
||||
scenario_ids = get_expert_scenario_ids(args.expert_data_path, max_ids=5)
|
||||
if not scenario_ids:
|
||||
print("[WARN] No scenario_id in expert pkl, falling back to scenarios [0,1,2]. Eval may have 0 controlled agents.")
|
||||
scenario_ids = [0, 1, 2]
|
||||
|
||||
total_rewards = []
|
||||
total_steps = []
|
||||
collision_episodes = 0
|
||||
horizon = 200
|
||||
single_agent = getattr(args, "single_agent", False)
|
||||
|
||||
for idx, scenario_id in enumerate(scenario_ids):
|
||||
env_config = {
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"horizon": horizon,
|
||||
"start_scenario_index": scenario_id,
|
||||
"num_scenarios": 1,
|
||||
"log_level": 50,
|
||||
}
|
||||
if single_agent:
|
||||
env = BCEgoReplayEnv(config=env_config)
|
||||
else:
|
||||
env = BCScenarioEnv(env_config, agent2policy=None)
|
||||
try:
|
||||
obs_dict = env.reset(seed=scenario_id)
|
||||
except Exception as e:
|
||||
print(f" Eval Episode {idx} (scenario {scenario_id}): reset failed: {e}")
|
||||
env.close()
|
||||
continue
|
||||
|
||||
n_controlled = len(env.controlled_agents)
|
||||
n_total_in_scenario = getattr(env, "num_controlled_in_scenario", n_controlled) if not single_agent else 1
|
||||
if n_controlled == 0:
|
||||
print(
|
||||
f" Eval Episode {idx} (scenario {scenario_id}): 0 controlled agents (total in scenario: {n_total_in_scenario}), skip."
|
||||
)
|
||||
env.close()
|
||||
continue
|
||||
|
||||
episode_reward = 0.0
|
||||
step_count = 0
|
||||
had_near_collision = False
|
||||
dones = {"__all__": False}
|
||||
while not dones["__all__"] and step_count < horizon:
|
||||
step_count += 1
|
||||
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, infos = env.step(action_dict)
|
||||
episode_reward += sum(rewards.values())
|
||||
if infos:
|
||||
for _aid, info in infos.items():
|
||||
if isinstance(info, dict) and info.get("near_collision", False):
|
||||
had_near_collision = True
|
||||
break
|
||||
|
||||
total_rewards.append(episode_reward)
|
||||
total_steps.append(step_count)
|
||||
if had_near_collision:
|
||||
collision_episodes += 1
|
||||
mode_str = "single-agent (ego)" if single_agent else f"agents (current): {n_controlled}, total in scenario: {n_total_in_scenario}"
|
||||
print(
|
||||
f" Eval Episode {idx} (scenario {scenario_id}): Total Reward {episode_reward:.2f}, steps {step_count}, {mode_str}"
|
||||
)
|
||||
env.close()
|
||||
|
||||
if not total_rewards:
|
||||
print(" No valid eval episodes (all skipped or failed).")
|
||||
return 0.0, 0.0, 0.0
|
||||
avg_reward = float(np.mean(total_rewards))
|
||||
avg_steps = float(np.mean(total_steps)) if total_steps else 0.0
|
||||
collision_rate = float(collision_episodes / max(1, len(total_rewards)))
|
||||
print(
|
||||
f" Average Evaluation Reward: {avg_reward:.2f} | Mean Episode Length: {avg_steps:.1f} | "
|
||||
f"Collision Rate (near): {collision_rate:.3f}"
|
||||
)
|
||||
return avg_reward, collision_rate, avg_steps
|
||||
|
||||
|
||||
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)
|
||||
|
||||
agent_id_filter = "default_agent" if getattr(args, "single_agent", False) else None
|
||||
obs_data, act_data = load_expert_pkl(
|
||||
args.expert_data_path,
|
||||
filter_terminal_last_step=args.filter_terminal_last_step,
|
||||
agent_id_filter=agent_id_filter,
|
||||
)
|
||||
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"))
|
||||
|
||||
# Periodic checkpointing (II-style)
|
||||
if args.checkpoint_freq > 0 and (epoch + 1) % args.checkpoint_freq == 0:
|
||||
torch.save(policy.state_dict(), os.path.join(args.save_dir, f"policy_epoch{epoch+1}.pt"))
|
||||
|
||||
if (epoch + 1) % args.eval_freq == 0:
|
||||
eval_reward, eval_collision_rate, eval_mean_steps = evaluate_policy(policy, args, device)
|
||||
writer.add_scalar("Reward/eval", eval_reward, epoch)
|
||||
writer.add_scalar("Eval/collision_rate_near", eval_collision_rate, epoch)
|
||||
writer.add_scalar("Eval/mean_episode_length", eval_mean_steps, 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)
|
||||
parser.add_argument("--checkpoint_freq", type=int, default=50, help="Save policy_epochN.pt every N epochs. Set <=0 to disable.")
|
||||
parser.add_argument(
|
||||
"--filter_terminal_last_step",
|
||||
action="store_true",
|
||||
help="Drop the last (obs, act) pair of each trajectory to approximate training on non-terminal steps (II-style).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--single_agent",
|
||||
action="store_true",
|
||||
help="Use single-agent (ego) expert data and evaluation; load only default_agent trajectories and evaluate with BCEgoReplayEnv.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
548
train_magail.py
Normal file
548
train_magail.py
Normal file
@@ -0,0 +1,548 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.optim as optim
|
||||
from torch.distributions import Normal
|
||||
import numpy as np
|
||||
import os
|
||||
import argparse
|
||||
import signal
|
||||
import sys
|
||||
from torch.utils.data import DataLoader
|
||||
from dataset.loader import MAGAILExpertDataset
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
|
||||
# --- Networks ---
|
||||
|
||||
class Actor(nn.Module):
|
||||
def __init__(self, state_dim, action_dim, hidden_dim=256):
|
||||
super(Actor, self).__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(state_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
nn.Linear(hidden_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
)
|
||||
self.mu_head = nn.Linear(hidden_dim, action_dim)
|
||||
self.log_std_head = nn.Parameter(torch.zeros(1, action_dim))
|
||||
|
||||
def forward(self, state):
|
||||
x = self.net(state)
|
||||
mu = torch.tanh(self.mu_head(x)) # Action range [-1, 1]
|
||||
if mu.dim() == 1:
|
||||
mu = mu.unsqueeze(0) # Handle single sample
|
||||
log_std = self.log_std_head.expand_as(mu)
|
||||
std = torch.exp(log_std)
|
||||
dist = Normal(mu, std)
|
||||
return dist
|
||||
|
||||
class Critic(nn.Module):
|
||||
def __init__(self, state_dim, hidden_dim=256):
|
||||
super(Critic, self).__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(state_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
nn.Linear(hidden_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
nn.Linear(hidden_dim, 1)
|
||||
)
|
||||
|
||||
def forward(self, state):
|
||||
return self.net(state)
|
||||
|
||||
class Discriminator(nn.Module):
|
||||
def __init__(self, state_dim, action_dim, hidden_dim=256):
|
||||
super(Discriminator, self).__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(state_dim + action_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
nn.Linear(hidden_dim, hidden_dim),
|
||||
nn.Tanh(),
|
||||
nn.Linear(hidden_dim, 1),
|
||||
nn.Sigmoid()
|
||||
)
|
||||
|
||||
def forward(self, state, action):
|
||||
x = torch.cat([state, action], dim=-1)
|
||||
return self.net(x)
|
||||
|
||||
# --- PPO Algorithm ---
|
||||
|
||||
class PPO:
|
||||
def __init__(self, state_dim, action_dim, lr=3e-4, gamma=0.99, eps_clip=0.2, K_epochs=10):
|
||||
self.actor = Actor(state_dim, action_dim).cuda()
|
||||
self.critic = Critic(state_dim).cuda()
|
||||
self.optimizer_actor = optim.Adam(self.actor.parameters(), lr=lr)
|
||||
self.optimizer_critic = optim.Adam(self.critic.parameters(), lr=lr)
|
||||
|
||||
self.gamma = gamma
|
||||
self.eps_clip = eps_clip
|
||||
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)
|
||||
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()
|
||||
dones = torch.FloatTensor(np.array(memory['dones'])).cuda()
|
||||
|
||||
# Monte Carlo estimate of state rewards (or GAE if implemented, simplistic here)
|
||||
# Usually for PPO we use GAE. Let's do a simple discounted return for now or bootstrapping.
|
||||
# Let's use bootstrapping from critic for returns.
|
||||
|
||||
returns = []
|
||||
discounted_reward = 0
|
||||
# This simple loop assumes full episode or consistent batch.
|
||||
# For multi-agent disjoint steps, bootstrapping is better.
|
||||
# But let's calculate advantage using GAE for stability.
|
||||
|
||||
values = self.critic(states).detach()
|
||||
next_values = self.critic(next_states).detach()
|
||||
|
||||
# GAE
|
||||
advantages = []
|
||||
gae = 0
|
||||
for i in reversed(range(len(rewards))):
|
||||
delta = rewards[i] + self.gamma * next_values[i] * (1 - dones[i]) - values[i]
|
||||
gae = delta + self.gamma * 0.95 * (1 - dones[i]) * gae
|
||||
advantages.insert(0, gae)
|
||||
|
||||
advantages = torch.FloatTensor(advantages).cuda()
|
||||
returns = advantages + values.squeeze()
|
||||
|
||||
# Optimize policy for K epochs:
|
||||
for _ in range(self.K_epochs):
|
||||
# Evaluating old actions and values :
|
||||
dist = self.actor(states)
|
||||
action_logprobs = self._log_prob_from_dist(dist, pre_tanh_actions)
|
||||
dist_entropy = dist.entropy().sum(dim=-1)
|
||||
state_values = self.critic(states).squeeze()
|
||||
|
||||
# Finding the ratio (pi_theta / pi_theta__old):
|
||||
ratios = torch.exp(action_logprobs - logprobs)
|
||||
|
||||
# Finding Surrogate Loss:
|
||||
surr1 = ratios * advantages
|
||||
surr2 = torch.clamp(ratios, 1-self.eps_clip, 1+self.eps_clip) * advantages
|
||||
loss = -torch.min(surr1, surr2) + 0.5*self.mse_loss(state_values, returns) - 0.01*dist_entropy
|
||||
|
||||
# take gradient step
|
||||
self.optimizer_actor.zero_grad()
|
||||
self.optimizer_critic.zero_grad()
|
||||
loss.mean().backward()
|
||||
self.optimizer_actor.step()
|
||||
self.optimizer_critic.step()
|
||||
|
||||
return loss.mean().item()
|
||||
|
||||
def save(self, checkpoint_path):
|
||||
torch.save(self.actor.state_dict(), checkpoint_path + "_actor.pth")
|
||||
torch.save(self.critic.state_dict(), checkpoint_path + "_critic.pth")
|
||||
|
||||
# --- Training Loop ---
|
||||
|
||||
def train(args):
|
||||
# 1. Setup Environment (45-dim obs via BCScenarioEnv)
|
||||
# Config for Env
|
||||
env_config = {
|
||||
"data_directory": args.data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3, # Dynamic
|
||||
"horizon": 200,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"start_scenario_index": 0,
|
||||
"num_scenarios": args.num_scenarios # Use argument
|
||||
}
|
||||
|
||||
# Ideally we use a wrapper for RL
|
||||
# env = MultiAgentScenarioEnv(config=env_config) # This requires Waymo data loader setup
|
||||
|
||||
# 2. Setup Models
|
||||
state_dim = 45
|
||||
action_dim = 2
|
||||
|
||||
ppo_agent = PPO(state_dim, action_dim)
|
||||
discriminator = Discriminator(state_dim, action_dim).cuda()
|
||||
disc_optimizer = optim.Adam(discriminator.parameters(), lr=3e-4)
|
||||
disc_criterion = nn.BCELoss()
|
||||
|
||||
# 3. Load Expert Data
|
||||
expert_dataset = MAGAILExpertDataset(args.expert_data_dir)
|
||||
# Ensure batch_size is not larger than dataset
|
||||
if len(expert_dataset) < args.batch_size:
|
||||
print(f"Warning: Expert dataset size {len(expert_dataset)} < batch_size {args.batch_size}. Adjusting batch_size.")
|
||||
args.batch_size = len(expert_dataset)
|
||||
if args.batch_size == 0:
|
||||
raise ValueError("Expert dataset is empty!")
|
||||
|
||||
expert_loader = DataLoader(expert_dataset, batch_size=args.batch_size, shuffle=True, drop_last=True)
|
||||
|
||||
# Create an infinite iterator
|
||||
def cycle(loader):
|
||||
while True:
|
||||
for batch in loader:
|
||||
yield batch
|
||||
expert_iter = cycle(expert_loader)
|
||||
|
||||
# 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?
|
||||
# But Env might return something else if we are using default ScenarioEnv settings.
|
||||
# ScenarioEnv returns list of obs.
|
||||
# The error says: "mat1 and mat2 shapes cannot be multiplied (1x108 and 45x256)"
|
||||
# This means the Env is returning 108-dim observation (MetaDrive default + Lidar),
|
||||
# but our Actor expects 45 (which is what we saved in expert data).
|
||||
|
||||
# We must align the environment observation space with our expert data format.
|
||||
# Our ExpertReplayEnv used a custom _get_all_obs.
|
||||
# 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 = BCScenarioEnv(env_config, agent2policy={}) # 45-dim obs
|
||||
|
||||
print("Starting training...")
|
||||
|
||||
# Tensorboard
|
||||
try:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
writer = SummaryWriter(log_dir=args.log_dir)
|
||||
except ImportError:
|
||||
print("TensorBoard not installed. Logging to console only.")
|
||||
writer = None
|
||||
|
||||
global_step = 0
|
||||
|
||||
for i_episode in range(args.max_episodes):
|
||||
# --- 1. Collect Rollouts (Interaction) ---
|
||||
memory = {
|
||||
'states': [],
|
||||
'actions': [],
|
||||
'pre_tanh_actions': [],
|
||||
'logprobs': [],
|
||||
'rewards': [],
|
||||
'next_states': [],
|
||||
'dones': []
|
||||
}
|
||||
|
||||
# Prepare seed
|
||||
available_scenarios = env.config["num_scenarios"]
|
||||
start_index = env.config["start_scenario_index"]
|
||||
seed = np.random.randint(start_index, start_index + available_scenarios)
|
||||
|
||||
# Reset Env
|
||||
try:
|
||||
# MetaDrive sometimes complains about uncleared objects if reset happens too fast or with lingering objs
|
||||
# We can try to force clear before reset or handle exception
|
||||
# But standard reset should handle it.
|
||||
# The error "You should clear all generated objects..." means some manager didn't clear its objects.
|
||||
# This is likely due to TrafficManager or AgentManager holding refs.
|
||||
|
||||
# Re-creating env is safer but slower.
|
||||
# Let's try closing and re-creating if reset fails frequently.
|
||||
# Or just ignore this error and try reset again? No, reset failing is fatal usually.
|
||||
|
||||
# Hack: Manually clear objects if we can access engine
|
||||
if env.engine is not None:
|
||||
env.engine.clear_objects(list(env.engine.get_objects().keys()))
|
||||
|
||||
obs_dict = env.reset(seed=seed)
|
||||
except Exception as e:
|
||||
# print(f"Env reset failed: {e}. Recreating environment...")
|
||||
try:
|
||||
env.close()
|
||||
except:
|
||||
pass
|
||||
|
||||
# Ensure engine is closed properly
|
||||
from metadrive.engine.engine_utils import close_engine
|
||||
try:
|
||||
close_engine()
|
||||
except Exception as e2:
|
||||
# Force cleanup of singleton if close failed
|
||||
from metadrive.engine.base_engine import BaseEngine
|
||||
if BaseEngine.singleton is not None:
|
||||
BaseEngine.singleton = None
|
||||
|
||||
# Also need to clear ShowBase
|
||||
try:
|
||||
from direct.showbase.ShowBase import ShowBase
|
||||
if hasattr(base, 'destroy'):
|
||||
base.destroy()
|
||||
except:
|
||||
pass
|
||||
|
||||
# Brutal force: delete base from builtins if it exists
|
||||
import builtins
|
||||
if hasattr(builtins, 'base'):
|
||||
del builtins.base
|
||||
|
||||
# print(f"Error closing engine: {e2}")
|
||||
|
||||
# Explicitly delete old env object to free memory
|
||||
del env
|
||||
import gc
|
||||
gc.collect()
|
||||
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
obs_dict = env.reset(seed=seed)
|
||||
|
||||
episode_reward = 0
|
||||
steps = 0
|
||||
|
||||
# Rollout loop
|
||||
while True:
|
||||
# Select actions for all agents
|
||||
actions = {}
|
||||
action_logprobs = {}
|
||||
pre_tanh_actions = {}
|
||||
|
||||
# obs_dict: {agent_id: obs}
|
||||
# MultiAgentScenarioEnv usually returns a dict {agent_id: obs}
|
||||
# BUT wait, check scenario_env.py implementation
|
||||
|
||||
if isinstance(obs_dict, list):
|
||||
# This happens if the environment returns a list instead of a dict
|
||||
# MultiAgentScenarioEnv._get_all_obs returns a list in original implementation?
|
||||
# Let's check scenario_env.py
|
||||
# If it returns list, we need to map it to agent ids or just iterate
|
||||
pass
|
||||
|
||||
# Temporary fix if it returns list (which means my previous edit to Env/expert_replay_env.py
|
||||
# changed it there, but maybe not in Env/scenario_env.py which we are using here!)
|
||||
|
||||
if isinstance(obs_dict, list):
|
||||
# We need agent IDs to step
|
||||
# In MultiAgentScenarioEnv, controlled_agents is a dict.
|
||||
# If obs is a list, it probably corresponds to controlled_agents.values() order?
|
||||
# This is risky.
|
||||
# Let's assume obs_dict is actually just observations.
|
||||
# We need to keys to create action dict.
|
||||
|
||||
current_agent_ids = list(env.controlled_agents.keys())
|
||||
# Ensure length matches
|
||||
if len(obs_dict) != len(current_agent_ids):
|
||||
# print(f"Warning: Obs list len {len(obs_dict)} != agents {len(current_agent_ids)}")
|
||||
pass
|
||||
|
||||
# Reconstruct dict
|
||||
new_obs_dict = {}
|
||||
for i, agent_id in enumerate(current_agent_ids):
|
||||
if i < len(obs_dict):
|
||||
new_obs_dict[agent_id] = obs_dict[i]
|
||||
obs_dict = new_obs_dict
|
||||
|
||||
for agent_id, obs in obs_dict.items():
|
||||
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)
|
||||
|
||||
# Store in memory
|
||||
for agent_id, obs in obs_dict.items():
|
||||
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)
|
||||
# For GAIL update we use Discriminator reward later
|
||||
memory['rewards'].append(0) # Placeholder
|
||||
|
||||
# Next state
|
||||
if agent_id in next_obs_dict:
|
||||
memory['next_states'].append(next_obs_dict[agent_id])
|
||||
memory['dones'].append(dones.get("__all__", False))
|
||||
else:
|
||||
# Agent finished/vanished
|
||||
# We need a dummy next state or handle done correctly
|
||||
# Just duplicate current state and mark done?
|
||||
memory['next_states'].append(obs)
|
||||
memory['dones'].append(True)
|
||||
|
||||
obs_dict = next_obs_dict
|
||||
steps += 1
|
||||
|
||||
if dones["__all__"] or steps >= 200: # Limit horizon
|
||||
break
|
||||
|
||||
# Initialize losses to 0/None before potential loop skip
|
||||
disc_loss = torch.tensor(0.0)
|
||||
ppo_loss = 0.0
|
||||
all_gail_rewards = [0.0]
|
||||
|
||||
# --- 2. Train Discriminator ---
|
||||
# Convert policy memory to tensors
|
||||
policy_states = torch.FloatTensor(np.array(memory['states'])).cuda()
|
||||
policy_actions = torch.FloatTensor(np.array(memory['actions'])).cuda()
|
||||
|
||||
# Sample expert batch
|
||||
expert_batch = next(expert_iter)
|
||||
|
||||
expert_states = expert_batch['state'].cuda()
|
||||
expert_actions = expert_batch['action'].cuda()
|
||||
|
||||
# Minibatch size matching
|
||||
batch_size = min(policy_states.size(0), expert_states.size(0))
|
||||
|
||||
if batch_size > 0: # Only train if we have data
|
||||
policy_states = policy_states[:batch_size]
|
||||
policy_actions = policy_actions[:batch_size]
|
||||
expert_states = expert_states[:batch_size]
|
||||
expert_actions = expert_actions[:batch_size]
|
||||
|
||||
# Update Discriminator
|
||||
# Label 1 for Expert, 0 for Policy
|
||||
# Train Expert
|
||||
disc_optimizer.zero_grad()
|
||||
|
||||
exp_preds = discriminator(expert_states, expert_actions)
|
||||
exp_loss = disc_criterion(exp_preds, torch.ones_like(exp_preds))
|
||||
|
||||
pol_preds = discriminator(policy_states.detach(), policy_actions.detach()) # Detach policy data
|
||||
pol_loss = disc_criterion(pol_preds, torch.zeros_like(pol_preds))
|
||||
|
||||
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))
|
||||
# Or more stable: log(D(s, a)) ? Original GAIL uses -log(1-D) which is log(D) roughly.
|
||||
# Let's use -log(1 - D(s, a) + eps)
|
||||
|
||||
# Actually PPO needs the full trajectory for GAE.
|
||||
# So we should compute rewards for ALL policy samples in memory.
|
||||
|
||||
all_policy_states = torch.FloatTensor(np.array(memory['states'])).cuda()
|
||||
all_policy_actions = torch.FloatTensor(np.array(memory['actions'])).cuda()
|
||||
|
||||
with torch.no_grad():
|
||||
all_d_val = discriminator(all_policy_states, all_policy_actions)
|
||||
all_gail_rewards = -torch.log(1 - all_d_val + 1e-8).cpu().numpy().flatten()
|
||||
|
||||
# Replace placeholders
|
||||
memory['rewards'] = all_gail_rewards.tolist()
|
||||
|
||||
# Update PPO
|
||||
ppo_loss = ppo_agent.update(memory)
|
||||
|
||||
# Clean up memory
|
||||
del policy_states, policy_actions, expert_states, expert_actions, exp_preds, exp_loss, pol_preds, pol_loss
|
||||
del all_policy_states, all_policy_actions, all_d_val
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
print(f"Episode {i_episode}: No data collected (Env might have crashed or no agents). Skipping update.")
|
||||
|
||||
# --- 4. Logging ---
|
||||
if writer:
|
||||
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.save_dir, f"model_{i_episode}"))
|
||||
|
||||
env.close()
|
||||
if writer:
|
||||
writer.close()
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--expert_data_dir", type=str, default="data/training_data", help="Directory with .pkl expert data")
|
||||
parser.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Waymo data dir for Env")
|
||||
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="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 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