Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 03dee0205a | |||
| 21c046aef0 | |||
| 265b0eade1 | |||
| 4dbea5f0a6 | |||
| c94571ddaa | |||
| 62e638c4d2 | |||
| b626702cbb | |||
| 22ce995916 |
57
.gitignore
vendored
Normal file
57
.gitignore
vendored
Normal file
@@ -0,0 +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
|
||||||
136
CHANGELOG.md
Normal file
136
CHANGELOG.md
Normal file
@@ -0,0 +1,136 @@
|
|||||||
|
# 更新日志
|
||||||
|
|
||||||
|
## 2025-01-20 问题修复与优化
|
||||||
|
|
||||||
|
### ✅ 已解决的问题
|
||||||
|
|
||||||
|
#### 1. 车辆生成位置偏差问题
|
||||||
|
**问题描述:** 部分车辆生成于草坪、停车场等非车道区域
|
||||||
|
|
||||||
|
**解决方案:**
|
||||||
|
- 实现 `_is_position_on_lane()` 方法:检测位置是否在有效车道上
|
||||||
|
- 实现 `_filter_valid_spawn_positions()` 方法:自动过滤非车道区域车辆
|
||||||
|
- 支持容差参数(默认3米)处理边界情况
|
||||||
|
- 在 `reset()` 时自动执行过滤,并输出统计信息
|
||||||
|
|
||||||
|
**配置参数:**
|
||||||
|
```python
|
||||||
|
"filter_offroad_vehicles": True, # 启用/禁用过滤
|
||||||
|
"lane_tolerance": 3.0, # 容差范围(米)
|
||||||
|
"max_controlled_vehicles": 10, # 最大车辆数限制
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 2. 红绿灯信息采集问题
|
||||||
|
**问题描述:**
|
||||||
|
- 部分红绿灯状态为 None
|
||||||
|
- 车道分段时部分车辆无法获取红绿灯状态
|
||||||
|
|
||||||
|
**解决方案:**
|
||||||
|
- 实现 `_get_traffic_light_state()` 方法,采用双重检测策略
|
||||||
|
- 方法1(优先):从导航模块获取当前车道,直接查询(高效)
|
||||||
|
- 方法2(兜底):遍历所有车道匹配位置(处理特殊情况)
|
||||||
|
- 完善异常处理,None 状态返回 0(无红绿灯)
|
||||||
|
- 返回值:0=无/未知, 1=绿灯, 2=黄灯, 3=红灯
|
||||||
|
|
||||||
|
#### 3. 性能优化问题
|
||||||
|
**问题描述:** FPS只有15帧,CPU利用率不到20%
|
||||||
|
|
||||||
|
**解决方案:**
|
||||||
|
- 创建 `run_multiagent_env_fast.py`:激光雷达优化版(30-60 FPS)
|
||||||
|
- 创建 `run_multiagent_env_parallel.py`:多进程并行版(300-600 steps/s)
|
||||||
|
- 提供详细的性能优化文档
|
||||||
|
|
||||||
|
### 📝 修改的文件
|
||||||
|
|
||||||
|
1. **Env/scenario_env.py**
|
||||||
|
- 新增 `_is_position_on_lane()` 方法
|
||||||
|
- 新增 `_filter_valid_spawn_positions()` 方法
|
||||||
|
- 新增 `_get_traffic_light_state()` 方法
|
||||||
|
- 更新 `default_config()` 添加配置参数
|
||||||
|
- 更新 `reset()` 调用过滤逻辑
|
||||||
|
- 更新 `_get_all_obs()` 使用新的红绿灯检测方法
|
||||||
|
|
||||||
|
2. **Env/run_multiagent_env.py**
|
||||||
|
- 添加车道过滤配置参数
|
||||||
|
|
||||||
|
3. **Env/run_multiagent_env_fast.py**
|
||||||
|
- 添加车道过滤配置
|
||||||
|
- 性能优化配置
|
||||||
|
|
||||||
|
4. **Env/run_multiagent_env_parallel.py**
|
||||||
|
- 添加车道过滤配置
|
||||||
|
- 多进程并行实现
|
||||||
|
|
||||||
|
5. **README.md**
|
||||||
|
- 更新问题说明,添加解决方案
|
||||||
|
- 添加配置示例和测试方法
|
||||||
|
- 添加问题解决总结
|
||||||
|
|
||||||
|
6. **新增文件**
|
||||||
|
- `Env/test_lane_filter.py`:功能测试脚本
|
||||||
|
|
||||||
|
### 🧪 测试方法
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 测试车道过滤和红绿灯检测功能
|
||||||
|
python Env/test_lane_filter.py
|
||||||
|
|
||||||
|
# 运行标准版本(带过滤和可视化)
|
||||||
|
python Env/run_multiagent_env.py
|
||||||
|
|
||||||
|
# 运行高性能版本(适合训练)
|
||||||
|
python Env/run_multiagent_env_fast.py
|
||||||
|
|
||||||
|
# 运行多进程并行版本(最高吞吐量)
|
||||||
|
python Env/run_multiagent_env_parallel.py
|
||||||
|
```
|
||||||
|
|
||||||
|
### 💡 使用建议
|
||||||
|
|
||||||
|
1. **调试阶段**:使用 `run_multiagent_env.py`,启用渲染和车道过滤
|
||||||
|
2. **训练阶段**:使用 `run_multiagent_env_fast.py`,关闭渲染,启用所有优化
|
||||||
|
3. **大规模训练**:使用 `run_multiagent_env_parallel.py`,充分利用多核CPU
|
||||||
|
|
||||||
|
### ⚙️ 配置说明
|
||||||
|
|
||||||
|
所有配置参数都可以在创建环境时通过 `config` 字典传递:
|
||||||
|
|
||||||
|
```python
|
||||||
|
env = MultiAgentScenarioEnv(
|
||||||
|
config={
|
||||||
|
# 基础配置
|
||||||
|
"data_directory": "...",
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"horizon": 300,
|
||||||
|
|
||||||
|
# 车道过滤(新增)
|
||||||
|
"filter_offroad_vehicles": True, # 启用车道过滤
|
||||||
|
"lane_tolerance": 3.0, # 容差3米
|
||||||
|
"max_controlled_vehicles": 10, # 最多10辆车
|
||||||
|
|
||||||
|
# 性能优化
|
||||||
|
"use_render": False,
|
||||||
|
"decision_repeat": 5,
|
||||||
|
...
|
||||||
|
},
|
||||||
|
agent2policy=your_policy
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 🔍 技术细节
|
||||||
|
|
||||||
|
**车道检测逻辑:**
|
||||||
|
1. 使用 `lane.lane.point_on_lane()` 精确检测
|
||||||
|
2. 使用 `lane.local_coordinates()` 计算横向距离
|
||||||
|
3. 支持容差参数处理边界情况
|
||||||
|
|
||||||
|
**红绿灯检测逻辑:**
|
||||||
|
1. 优先从 `vehicle.navigation.current_lane` 获取
|
||||||
|
2. 失败时遍历所有车道查找
|
||||||
|
3. 所有异常均有保护,确保稳定性
|
||||||
|
|
||||||
|
**性能优化原理:**
|
||||||
|
- 减少激光束数量降低计算量
|
||||||
|
- 多进程绕过Python GIL限制
|
||||||
|
- 充分利用多核CPU
|
||||||
|
|
||||||
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.
BIN
Env/__pycache__/logger_utils.cpython-310.pyc
Normal file
BIN
Env/__pycache__/logger_utils.cpython-310.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/logger_utils.cpython-313.pyc
Normal file
BIN
Env/__pycache__/logger_utils.cpython-313.pyc
Normal file
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.
BIN
Env/__pycache__/run_multiagent_env.cpython-310.pyc
Normal file
BIN
Env/__pycache__/run_multiagent_env.cpython-310.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/scenario_env.cpython-310.pyc
Normal file
BIN
Env/__pycache__/scenario_env.cpython-310.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/scenario_env.cpython-313.pyc
Normal file
BIN
Env/__pycache__/scenario_env.cpython-313.pyc
Normal file
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.
BIN
Env/__pycache__/simple_idm_policy.cpython-310.pyc
Normal file
BIN
Env/__pycache__/simple_idm_policy.cpython-310.pyc
Normal file
Binary file not shown.
BIN
Env/__pycache__/simple_idm_policy.cpython-313.pyc
Normal file
BIN
Env/__pycache__/simple_idm_policy.cpython-313.pyc
Normal file
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.
64
Env/bc_env.py
Normal file
64
Env/bc_env.py
Normal file
@@ -0,0 +1,64 @@
|
|||||||
|
from Env.scenario_env import MultiAgentScenarioEnv
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
class BCScenarioEnv(MultiAgentScenarioEnv):
|
||||||
|
"""
|
||||||
|
Environment for Behavior Cloning Evaluation.
|
||||||
|
Uses the same 45-dim observation as ExpertReplayEnv:
|
||||||
|
- Ego State (5): x, y, vx, vy, heading
|
||||||
|
- Neighbors (40): 10 nearest * (rel_x, rel_y, vx, vy)
|
||||||
|
"""
|
||||||
|
def _get_all_obs(self):
|
||||||
|
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||||
|
obs_dict = {}
|
||||||
|
|
||||||
|
for agent_id, vehicle in self.controlled_agents.items():
|
||||||
|
# 1. Ego State
|
||||||
|
ego_state = [
|
||||||
|
vehicle.position[0], vehicle.position[1],
|
||||||
|
vehicle.velocity[0], vehicle.velocity[1],
|
||||||
|
vehicle.heading_theta
|
||||||
|
]
|
||||||
|
|
||||||
|
# 2. Neighbors
|
||||||
|
neighbors = []
|
||||||
|
# Iterate through all vehicles in the engine
|
||||||
|
candidates = []
|
||||||
|
# Use engine.agent_manager.active_agents to find neighbors
|
||||||
|
# Note: This includes background vehicles if they are in active_agents
|
||||||
|
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||||
|
if other_id == agent_id:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Check if vehicle is valid/active
|
||||||
|
# (MetaDrive manages active_agents, so they should be active)
|
||||||
|
|
||||||
|
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||||
|
if dist < 30.0:
|
||||||
|
candidates.append((dist, other_vehicle))
|
||||||
|
|
||||||
|
# Sort by distance
|
||||||
|
candidates.sort(key=lambda x: x[0])
|
||||||
|
|
||||||
|
# Take top 10
|
||||||
|
top_10 = candidates[:10]
|
||||||
|
|
||||||
|
neighbor_feats = []
|
||||||
|
for _, neighbor in top_10:
|
||||||
|
neighbor_feats.extend([
|
||||||
|
neighbor.position[0] - vehicle.position[0], # Relative pos
|
||||||
|
neighbor.position[1] - vehicle.position[1],
|
||||||
|
neighbor.velocity[0], # Absolute vel
|
||||||
|
neighbor.velocity[1]
|
||||||
|
])
|
||||||
|
|
||||||
|
# Pad if < 10
|
||||||
|
missing = 10 - len(top_10)
|
||||||
|
if missing > 0:
|
||||||
|
neighbor_feats.extend([0.0] * (4 * missing))
|
||||||
|
|
||||||
|
# Flatten
|
||||||
|
obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||||
|
obs_dict[agent_id] = obs
|
||||||
|
|
||||||
|
return obs_dict
|
||||||
527
Env/expert_replay_env.py
Normal file
527
Env/expert_replay_env.py
Normal file
@@ -0,0 +1,527 @@
|
|||||||
|
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 = {} # Vehicles that exist but are static/background
|
||||||
|
|
||||||
|
# Helper function to check if a position is on a valid lane
|
||||||
|
def is_on_lane(pos, map_manager, threshold=2.0):
|
||||||
|
# Check if point is close to any lane in the road network
|
||||||
|
# This can be expensive if checked for every point, so we check sample points
|
||||||
|
# or rely on lane index if available.
|
||||||
|
# Waymo tracks don't have lane index, just positions.
|
||||||
|
# We can use map.road_network.get_closest_lane_index(pos)
|
||||||
|
if map_manager is None or map_manager.current_map is None:
|
||||||
|
return True # If no map, assume valid
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Use a larger search radius to catch slightly offset lanes
|
||||||
|
lane, lane_index = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True)
|
||||||
|
if lane is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Check lateral distance
|
||||||
|
long, lat = lane.local_coordinates(pos)
|
||||||
|
width = lane.width
|
||||||
|
# Allow being slightly off-lane (e.g. changing lanes)
|
||||||
|
# But parking lots are usually far from defined lanes in Waymo converted maps
|
||||||
|
if abs(lat) <= (width / 2 + threshold):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
except:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# --- MODIFIED SECTION START ---
|
||||||
|
# Capture expert tracks before they are cleaned
|
||||||
|
self.expert_tracks = {}
|
||||||
|
# Capture SDC track for ego replay (MetaDrive default agent)
|
||||||
|
self.sdc_track = None
|
||||||
|
self.sdc_vehicle = 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)
|
||||||
|
_obj_to_clean_this_frame = []
|
||||||
|
self.car_birth_info_list = []
|
||||||
|
|
||||||
|
# Pre-filter: Check tracks against map AND check for static vehicles
|
||||||
|
|
||||||
|
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']
|
||||||
|
if not valid.any():
|
||||||
|
continue
|
||||||
|
|
||||||
|
first_show = np.argmax(valid)
|
||||||
|
last_show = len(valid) - 1 - np.argmax(valid[::-1])
|
||||||
|
mid_show = (first_show + last_show) // 2
|
||||||
|
|
||||||
|
# 1. Lane check (existing logic)
|
||||||
|
points_to_check = [first_show, mid_show, last_show]
|
||||||
|
on_road_count = 0
|
||||||
|
is_valid_track = True
|
||||||
|
start_pos = track['state']['position'][first_show]
|
||||||
|
if not is_on_lane(start_pos, self.engine.map_manager, threshold=5.0): # 5m tolerance
|
||||||
|
mid_pos = track['state']['position'][mid_show]
|
||||||
|
if not is_on_lane(mid_pos, self.engine.map_manager, threshold=5.0):
|
||||||
|
is_valid_track = False
|
||||||
|
|
||||||
|
# 2. Static check
|
||||||
|
# Calculate total displacement and max speed
|
||||||
|
positions = track['state']['position'][valid.astype(bool)]
|
||||||
|
velocities = track['state']['velocity'][valid.astype(bool)]
|
||||||
|
|
||||||
|
total_displacement = 0
|
||||||
|
max_speed = 0
|
||||||
|
if len(positions) > 1:
|
||||||
|
total_displacement = np.linalg.norm(positions[-1] - positions[0])
|
||||||
|
max_speed = np.max(np.linalg.norm(velocities, axis=1))
|
||||||
|
|
||||||
|
is_static = False
|
||||||
|
if total_displacement < 5.0 and max_speed < 1.0: # Relaxed threshold: <5m move and <1m/s
|
||||||
|
is_static = True
|
||||||
|
|
||||||
|
# Decision logic:
|
||||||
|
# - If off-road AND static: Skip completely (don't even spawn as background)
|
||||||
|
# - If off-road but moving: Maybe keep? Or skip? Usually off-road moving is weird, skip.
|
||||||
|
# - If on-road but static: Spawn as BACKGROUND (visible but not controlled agent)
|
||||||
|
# - If on-road and moving: Spawn as CONTROLLED agent
|
||||||
|
|
||||||
|
if not is_valid_track:
|
||||||
|
# Skip off-road vehicles entirely (both static and moving off-road)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if is_static:
|
||||||
|
# Add to background list, but NOT to car_birth_info_list (which is for controlled agents)
|
||||||
|
# We need a way to spawn them. Let's add a separate list.
|
||||||
|
self.background_vehicles[scenario_id] = {
|
||||||
|
'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]),
|
||||||
|
'scenario_id': scenario_id,
|
||||||
|
'length': track['state']['length'][first_show],
|
||||||
|
'width': track['state']['width'][first_show],
|
||||||
|
'valid': valid # Need validity to know when to show/hide
|
||||||
|
}
|
||||||
|
continue # Do not add to controlled list
|
||||||
|
|
||||||
|
# Store the full track for replay (only for controlled agents)
|
||||||
|
self.expert_tracks[scenario_id] = track
|
||||||
|
|
||||||
|
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]),
|
||||||
|
'scenario_id': scenario_id, # Keep track of original ID to lookup tracks
|
||||||
|
'length': track['state']['length'][first_show],
|
||||||
|
'width': track['state']['width'][first_show]
|
||||||
|
})
|
||||||
|
|
||||||
|
for scenario_id in _obj_to_clean_this_frame:
|
||||||
|
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||||
|
# --- MODIFIED SECTION END ---
|
||||||
|
|
||||||
|
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_background_vehicles() # Initial spawn for background
|
||||||
|
|
||||||
|
# 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_background_vehicles(self):
|
||||||
|
# Spawn static/background vehicles
|
||||||
|
# Since they are static, we might just spawn them once if their show_time is 0
|
||||||
|
# But Waymo tracks have valid bits, they might appear/disappear.
|
||||||
|
# For optimization, if they are truly static (never move), we just spawn them when show_time matches.
|
||||||
|
|
||||||
|
# We need to track spawned background vehicles to remove them if they become invalid?
|
||||||
|
# Since we defined them as "static", they probably stay put.
|
||||||
|
# But validity might change (e.g. late spawn).
|
||||||
|
|
||||||
|
# For simplicity in this step, let's just iterate and spawn if time matches
|
||||||
|
for sid, car in self.background_vehicles.items():
|
||||||
|
if car['show_time'] == self.round:
|
||||||
|
# Spawn as a Traffic Vehicle (not PolicyVehicle), or just a static object?
|
||||||
|
# Using DefaultVehicle is fine, but don't add to controlled_agents
|
||||||
|
|
||||||
|
# Check duplication
|
||||||
|
bg_id = f"bg_{car['id']}"
|
||||||
|
# if bg_id in self.engine.obj_to_id: # obj_to_id might not be available in all versions
|
||||||
|
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']
|
||||||
|
)
|
||||||
|
|
||||||
|
# Set color to grey/dark to indicate background
|
||||||
|
v.set_velocity([0, 0])
|
||||||
|
# Maybe set color? MetaDrive vehicles random color.
|
||||||
|
# v.set_color(...) if supported
|
||||||
|
|
||||||
|
# Register as an active object but NOT controlled agent
|
||||||
|
# The engine manages it.
|
||||||
|
# CRITICAL: We need it in self.engine.agent_manager.active_agents for Observation?
|
||||||
|
# If we want it to be seen by Lidar/Observation, it needs to be an "agent" or "traffic".
|
||||||
|
# DefaultVehicle spawned this way is just an object.
|
||||||
|
# We should add it to traffic manager? Or just leave it as object?
|
||||||
|
# MultiAgentScenarioEnv._get_all_obs iterates self.engine.agent_manager.active_agents
|
||||||
|
|
||||||
|
# If we want it in observation, we must add it to active_agents OR iterate over all objects.
|
||||||
|
# Adding to active_agents is easier for compatibility.
|
||||||
|
self.engine.agent_manager.active_agents[bg_id] = v
|
||||||
|
|
||||||
|
# Store valid mask to remove it later if needed?
|
||||||
|
v.valid_mask = car['valid']
|
||||||
|
v.start_t = car['show_time']
|
||||||
|
|
||||||
|
def _update_background_vehicles(self):
|
||||||
|
# Remove background vehicles if they become invalid
|
||||||
|
# Or spawn new ones
|
||||||
|
self._spawn_background_vehicles()
|
||||||
|
|
||||||
|
# Check validity for existing
|
||||||
|
to_remove = []
|
||||||
|
for aid, v in self.engine.agent_manager.active_agents.items():
|
||||||
|
if aid.startswith("bg_"):
|
||||||
|
# Check validity
|
||||||
|
if hasattr(v, 'valid_mask'):
|
||||||
|
curr_step = self.round
|
||||||
|
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
||||||
|
to_remove.append(aid)
|
||||||
|
|
||||||
|
for aid in to_remove:
|
||||||
|
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||||
|
# if aid in self.engine.obj_to_id:
|
||||||
|
# self.engine.clear_objects([self.engine.obj_to_id[aid]])
|
||||||
|
# Instead, we should find the object by ID and clear it.
|
||||||
|
# Since we don't track obj directly, we can't easily clear it without obj ref.
|
||||||
|
# Wait, active_agents stores the vehicle object.
|
||||||
|
# So we can just clear that object.
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Re-iterate to clear objects properly
|
||||||
|
for aid in to_remove:
|
||||||
|
# We need to find the vehicle object to clear it.
|
||||||
|
# But we popped it from active_agents.
|
||||||
|
# Wait, we should get it before pop.
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _update_background_vehicles(self):
|
||||||
|
# Remove background vehicles if they become invalid
|
||||||
|
# Or spawn new ones
|
||||||
|
self._spawn_background_vehicles()
|
||||||
|
|
||||||
|
# Check validity for existing
|
||||||
|
to_remove = []
|
||||||
|
objects_to_clear = []
|
||||||
|
|
||||||
|
for aid, v in self.engine.agent_manager.active_agents.items():
|
||||||
|
if aid.startswith("bg_"):
|
||||||
|
# Check validity
|
||||||
|
if hasattr(v, 'valid_mask'):
|
||||||
|
curr_step = self.round
|
||||||
|
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
||||||
|
to_remove.append(aid)
|
||||||
|
objects_to_clear.append(v)
|
||||||
|
|
||||||
|
for aid in to_remove:
|
||||||
|
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||||
|
|
||||||
|
if objects_to_clear:
|
||||||
|
self.engine.clear_objects(objects_to_clear)
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
rewards = {aid: 0.0 for aid in self.controlled_agents}
|
||||||
|
dones = {aid: False for aid in self.controlled_agents}
|
||||||
|
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 self.controlled_agents}
|
||||||
|
|
||||||
|
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 = []
|
||||||
|
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))
|
||||||
|
|
||||||
|
# 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? or Relative? Usually relative in MultiAgent
|
||||||
|
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
|
||||||
65
Env/inverse_dynamics.py
Normal file
65
Env/inverse_dynamics.py
Normal file
@@ -0,0 +1,65 @@
|
|||||||
|
import numpy as np
|
||||||
|
import math
|
||||||
|
|
||||||
|
class InverseDynamics:
|
||||||
|
def __init__(self, max_steering=0.7, max_acc=15.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}
|
||||||
@@ -2,7 +2,7 @@ from scenario_env import MultiAgentScenarioEnv
|
|||||||
from Env.simple_idm_policy import ConstantVelocityPolicy
|
from Env.simple_idm_policy import ConstantVelocityPolicy
|
||||||
from metadrive.engine.asset_loader import AssetLoader
|
from metadrive.engine.asset_loader import AssetLoader
|
||||||
|
|
||||||
WAYMO_DATA_DIR = r"/home/zhy/桌面/MAGAIL_TR/Env"
|
WAYMO_DATA_DIR = r"/home/huangfukk/MAGAIL4AutoDrive/data"
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
env = MultiAgentScenarioEnv(
|
env = MultiAgentScenarioEnv(
|
||||||
|
|||||||
@@ -100,6 +100,13 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
for scenario_id in _obj_to_clean_this_frame:
|
for scenario_id in _obj_to_clean_this_frame:
|
||||||
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
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.engine.reset()
|
||||||
self.reset_sensors()
|
self.reset_sensors()
|
||||||
self.engine.taskMgr.step()
|
self.engine.taskMgr.step()
|
||||||
@@ -114,9 +121,6 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
self.episode_rewards = defaultdict(float)
|
self.episode_rewards = defaultdict(float)
|
||||||
self.episode_lengths = defaultdict(int)
|
self.episode_lengths = defaultdict(int)
|
||||||
|
|
||||||
self.controlled_agents.clear()
|
|
||||||
self.controlled_agent_ids.clear()
|
|
||||||
|
|
||||||
super().reset(seed) # 初始化场景
|
super().reset(seed) # 初始化场景
|
||||||
self._spawn_controlled_agents()
|
self._spawn_controlled_agents()
|
||||||
|
|
||||||
@@ -190,6 +194,7 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
self.controlled_agents[agent_id].before_step(action)
|
self.controlled_agents[agent_id].before_step(action)
|
||||||
|
|
||||||
self.engine.step()
|
self.engine.step()
|
||||||
|
self.engine.after_step()
|
||||||
|
|
||||||
for agent_id in action_dict:
|
for agent_id in action_dict:
|
||||||
if agent_id in self.controlled_agents:
|
if agent_id in self.controlled_agents:
|
||||||
|
|||||||
125
README.md
125
README.md
@@ -1,28 +1,121 @@
|
|||||||
# MAGAIL4AutoDrive
|
# MAGAIL4AutoDrive
|
||||||
### 1.1 环境搭建
|
|
||||||
环境核心代码封装于`Env`文件夹,通过运行`run_multiagent_env.py`即可启动多智能体交互环境,该脚本的核心功能为读取各智能体(车辆)的动作指令,并将其传入`env.step()`方法中完成仿真执行。
|
|
||||||
|
|
||||||
当前已初步实现`Env.senario_env.MultiAgentScenarioEnv.reset()`车辆生成函数,具体逻辑如下:首先读取专家数据集中各车辆的初始位姿信息;随后对原始数据进行清洗,剔除车辆 Agent 实例信息,记录核心参数(车辆 ID、初始生成位置、朝向角、生成时间戳、目标终点坐标);最后调用`_spawn_controlled_agents()`函数,依据清洗后的参数在指定时间、指定位置生成搭载自动驾驶算法的可控车辆。
|
基于 **MetaDrive** 仿真器和 **Waymo Open Motion Dataset** 的自动驾驶多智能体模仿学习(MAGAIL)与行为克隆(BC)训练系统。
|
||||||
|
|
||||||
需解决的关键问题:部分车辆存在生成位置偏差(如生成于草坪区域),推测成因可能为专家数据记录误差或场景中模拟停车场区域的特殊标注。后续计划引入车道区域检测机制,通过判断车辆初始生成位置是否位于有效车道范围内,对非车道区域生成的车辆进行过滤,确保环境初始化的合理性。
|
本项目旨在从真实的 Waymo 驾驶数据中提取专家轨迹,并通过模仿学习(Imitation Learning)训练能够适应复杂交互场景的自动驾驶策略。
|
||||||
|
|
||||||
|
## 目录结构
|
||||||
|
|
||||||
### 1.2 观测获取
|
```text
|
||||||
观测信息采集功能通过`Env.senario_env.MultiAgentScenarioEnv._get_all_obs()`函数实现,该函数支持遍历所有可控车辆并采集多维度观测数据,当前已实现的观测维度包括:车辆实时位置坐标、朝向角、行驶速度、雷达扫描点云(含障碍物与车道线特征)、导航信息(因场景复杂度较低,暂采用目标终点坐标直接作为导航输入)。
|
MAGAIL4AutoDrive/
|
||||||
|
├── Algorithm/ # 强化学习与模仿学习算法实现
|
||||||
|
│ ├── policy.py # 基础策略网络 (MLP 等)
|
||||||
|
│ ├── ppo.py # PPO 算法实现
|
||||||
|
│ ├── magail.py # MAGAIL 算法核心逻辑
|
||||||
|
│ ├── disc.py # 判别器 (Discriminator) 网络
|
||||||
|
│ └── ...
|
||||||
|
├── Env/ # 仿真环境封装 (MetaDrive Wrapper)
|
||||||
|
│ ├── bc_env.py # BCScenarioEnv,45 维观测(BC/MAGAIL 共用)
|
||||||
|
│ ├── scenario_env.py # 多智能体基础场景环境
|
||||||
|
│ ├── expert_replay_env.py # 专家轨迹回放环境(数据生成与回放)
|
||||||
|
│ ├── inverse_dynamics.py # 逆动力学模块 (轨迹 -> 动作)
|
||||||
|
│ ├── simple_idm_policy.py # ConstantVelocityPolicy 占位策略
|
||||||
|
│ └── ...
|
||||||
|
├── dataset/ # 数据集加载器
|
||||||
|
│ ├── loader.py # 主流水线:load_expert_pkl、MAGAILExpertDataset
|
||||||
|
│ └── expert_dataset.py # 可选 107 维/5 维管线
|
||||||
|
├── scripts/ # 工具脚本(数据、回放、可视化、分析)
|
||||||
|
│ ├── generate_expert_data.py # 从 Waymo 生成专家 (obs, act) pkl
|
||||||
|
│ ├── visualize.py # 可视化统一入口(replay / policy / trajectory)
|
||||||
|
│ ├── analyze_expert_data.py # 数据分布分析
|
||||||
|
│ ├── launch_tensorboard.py # 启动 TensorBoard
|
||||||
|
│ ├── README.md # 脚本用法说明
|
||||||
|
│ └── ...
|
||||||
|
├── data/ # 数据目录(相对路径)
|
||||||
|
│ ├── exp_filtered/ # Waymo 场景数据
|
||||||
|
│ ├── training_data/ # 专家 pkl 输出(generate_expert_data)
|
||||||
|
│ └── trajectories/ # 其他轨迹 pkl(如 expert_dataset 输出)
|
||||||
|
├── models/ # 模型保存目录(相对路径)
|
||||||
|
│ ├── bc/ # BC 模型 (.pt)
|
||||||
|
│ └── magail/ # MAGAIL 模型 (*_actor.pth, *_critic.pth)
|
||||||
|
├── logs/ # 训练日志 (TensorBoard)
|
||||||
|
│ ├── bc/
|
||||||
|
│ └── magail/
|
||||||
|
├── train_bc.py # [根目录] BC 训练
|
||||||
|
├── train_magail.py # [根目录] MAGAIL 训练
|
||||||
|
└── README.md
|
||||||
|
```
|
||||||
|
|
||||||
红绿灯信息采集机制需改进:当前方案通过 “车辆所属车道序号匹配对应红绿灯实例” 的逻辑获取信号灯状态,但存在两类问题:一是部分红绿灯实例的状态值为`None`;二是当单条车道存在分段设计时,部分区域的车辆会无法获取红绿灯状态。
|
## 路径约定(相对项目根)
|
||||||
|
|
||||||
|
- **数据**: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.4 动作执行
|
**1) 下载 Waymo Motion(TFRecord)**
|
||||||
在当前环境测试阶段,暂沿用腾达的动作执行框架:为每辆可控车辆分配独立的`policy`模型,将单车辆观测数据输入对应`policy`得到动作指令后,传入`env.step()`完成仿真;同时在`before_step`阶段调用`_set_action()`函数,将动作指令绑定至车辆实例,最终由 MetaDrive 仿真系统完成物理动力学计算与场景渲染。
|
安装 `gsutil` 并登录 Google 账号后,例如只下载 training_20s:
|
||||||
|
|
||||||
后续优化方向为构建 “参数共享式统一模型框架”,具体设计如下:所有车辆共用 1 个`policy`模型,通过参数共享机制实现模型的全局统一维护。该框架具备三重优势:一是避免多车辆独立模型带来的训练偏差(如不同模型训练程度不一致);二是解决车辆数量动态变化时的模型管理问题(车辆新增无需额外初始化模型,车辆减少不丢失模型训练信息);三是支持动作指令的并行计算,可显著提升每一步决策的迭代效率,适配大规模多智能体交互场景的训练需求。
|
```bash
|
||||||
|
gsutil -m cp -r "gs://waymo_open_dataset_motion_v_1_2_0/uncompressed/scenario/training_20s" ./waymo/
|
||||||
|
```
|
||||||
|
|
||||||
|
**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`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||||
|
```
|
||||||
|
|
||||||
|
## 核心工作流
|
||||||
|
|
||||||
|
### 1. 数据准备
|
||||||
|
使用 `scripts/generate_expert_data.py` 将 Waymo 数据转换为训练用 `.pkl`,输出到 `data/training_data/`。
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. 行为克隆 (BC)
|
||||||
|
- **训练**:`python train_bc.py`(模型保存到 `models/bc/`,日志到 `logs/bc/`)
|
||||||
|
- **可视化**:`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt`
|
||||||
|
|
||||||
|
### 3. 多智能体对抗模仿学习 (MAGAIL)
|
||||||
|
- **训练**:`python train_magail.py`(模型保存到 `models/magail/`,日志到 `logs/magail/`)
|
||||||
|
- **可视化**:`python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth`
|
||||||
|
|
||||||
|
### 4. 可视化统一入口
|
||||||
|
可视化统一使用 `scripts/visualize.py`,子命令:`replay`(场景回放)、`policy`(BC/MAGAIL 策略)、`trajectory`(专家轨迹 2D 动画)。详见 [scripts/README.md](scripts/README.md)。
|
||||||
|
|
||||||
|
## 文件与模块职责
|
||||||
|
|
||||||
|
### 根目录脚本
|
||||||
|
- **train_bc.py**:BC 训练,从 `dataset.loader` 加载专家 pkl,模型与日志写入 `models/bc/`、`logs/bc/`
|
||||||
|
- **train_magail.py**:MAGAIL 训练,环境使用 `BCScenarioEnv`(45 维),从 `dataset.loader` 加载专家数据,模型与日志写入 `models/magail/`、`logs/magail/`
|
||||||
|
|
||||||
|
### Env 模块
|
||||||
|
- **Env/bc_env.py**:`BCScenarioEnv`,45 维观测(Ego 5 维 + 10 邻居×4 维),BC 与 MAGAIL 训练/评估共用
|
||||||
|
- **Env/scenario_env.py**:`MultiAgentScenarioEnv` 基类,Waymo 场景加载与步进
|
||||||
|
- **Env/expert_replay_env.py**:专家轨迹回放与逆动力学动作,供 `generate_expert_data.py` 与回放可视化
|
||||||
|
- **Env/inverse_dynamics.py**:轨迹 → 油门/转向动作
|
||||||
|
|
||||||
|
### Algorithm 模块
|
||||||
|
- **Algorithm/policy.py**:`StateIndependentPolicy`,BC 使用的 MLP 策略
|
||||||
|
|
||||||
|
### scripts 目录
|
||||||
|
工具脚本用途与用法见 [scripts/README.md](scripts/README.md)。
|
||||||
|
|||||||
498
TRAINING_ARCHITECTURE.md
Normal file
498
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 作为策略优化算法。当前主要问题集中在环境稳定性和日志记录方面,需要进一步优化。
|
||||||
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("(使用之前已成功的方法)")
|
||||||
103
dataset/loader.py
Normal file
103
dataset/loader.py
Normal file
@@ -0,0 +1,103 @@
|
|||||||
|
"""
|
||||||
|
统一数据加载:BC/MAGAIL 训练用专家 pkl 的加载函数与 Dataset。
|
||||||
|
主训练流水线使用本模块;dataset/expert_dataset.py 为可选 107 维/5 维管线。
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import glob
|
||||||
|
import pickle
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
|
||||||
|
def load_expert_pkl(expert_data_path):
|
||||||
|
"""从目录或单个 pkl 加载专家 (obs, acts),返回 concat 后的 obs_data, act_data。"""
|
||||||
|
if os.path.isdir(expert_data_path):
|
||||||
|
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
||||||
|
if not pkl_files:
|
||||||
|
raise FileNotFoundError(f"No .pkl files in {expert_data_path}")
|
||||||
|
print(f"Found {len(pkl_files)} pickle files in {expert_data_path}")
|
||||||
|
elif os.path.exists(expert_data_path):
|
||||||
|
pkl_files = [expert_data_path]
|
||||||
|
else:
|
||||||
|
raise FileNotFoundError(f"Expert data path not found: {expert_data_path}")
|
||||||
|
|
||||||
|
obs_data, act_data = [], []
|
||||||
|
for pkl_file in pkl_files:
|
||||||
|
try:
|
||||||
|
with open(pkl_file, "rb") as f:
|
||||||
|
data = pickle.load(f)
|
||||||
|
if isinstance(data, list):
|
||||||
|
for traj in data:
|
||||||
|
if "obs" in traj and "acts" in traj:
|
||||||
|
obs_data.append(traj["obs"])
|
||||||
|
act_data.append(traj["acts"])
|
||||||
|
elif isinstance(data, dict):
|
||||||
|
if "observations" in data and "actions" in data:
|
||||||
|
obs_data.append(data["observations"])
|
||||||
|
act_data.append(data["actions"])
|
||||||
|
else:
|
||||||
|
print(f"Skipping {pkl_file}: Unknown data format {type(data)}")
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error loading {pkl_file}: {e}")
|
||||||
|
|
||||||
|
if len(obs_data) == 0:
|
||||||
|
raise ValueError("No valid data loaded from provided path.")
|
||||||
|
obs_data = np.concatenate(obs_data, axis=0)
|
||||||
|
act_data = np.concatenate(act_data, axis=0)
|
||||||
|
print(f"Total loaded samples: {len(obs_data)}")
|
||||||
|
return obs_data, act_data
|
||||||
|
|
||||||
|
|
||||||
|
class MAGAILExpertDataset(Dataset):
|
||||||
|
def __init__(self, data_dir, transform=None):
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
data_dir (str): Directory containing .pkl files from generate_expert_data.py
|
||||||
|
transform (callable, optional): Optional transform to be applied on a sample.
|
||||||
|
"""
|
||||||
|
self.data_dir = data_dir
|
||||||
|
self.transform = transform
|
||||||
|
self.trajectories = []
|
||||||
|
self.flat_data = [] # (obs, act) pairs
|
||||||
|
|
||||||
|
# Load all .pkl files
|
||||||
|
pkl_files = glob.glob(os.path.join(data_dir, "*.pkl"))
|
||||||
|
print(f"Loading data from {len(pkl_files)} files in {data_dir}...")
|
||||||
|
|
||||||
|
for pkl_file in pkl_files:
|
||||||
|
try:
|
||||||
|
with open(pkl_file, "rb") as f:
|
||||||
|
data = pickle.load(f)
|
||||||
|
# data is a list of dicts: {'obs': (T, 45), 'acts': (T, 2), ...}
|
||||||
|
self.trajectories.extend(data)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error loading {pkl_file}: {e}")
|
||||||
|
|
||||||
|
# Flatten for training Discriminator/BC
|
||||||
|
print(f"Processing {len(self.trajectories)} trajectories...")
|
||||||
|
for traj in self.trajectories:
|
||||||
|
obs = traj["obs"]
|
||||||
|
acts = traj["acts"]
|
||||||
|
|
||||||
|
# obs: (T, 45), acts: (T, 2)
|
||||||
|
for i in range(len(obs)):
|
||||||
|
self.flat_data.append((obs[i], acts[i]))
|
||||||
|
|
||||||
|
print(f"Total samples: {len(self.flat_data)}")
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.flat_data)
|
||||||
|
|
||||||
|
def __getitem__(self, idx):
|
||||||
|
obs, act = self.flat_data[idx]
|
||||||
|
|
||||||
|
obs = torch.from_numpy(obs).float()
|
||||||
|
act = torch.from_numpy(act).float()
|
||||||
|
|
||||||
|
sample = {"state": obs, "action": act}
|
||||||
|
|
||||||
|
if self.transform:
|
||||||
|
sample = self.transform(sample)
|
||||||
|
|
||||||
|
return sample
|
||||||
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.
76
scripts/README.md
Normal file
76
scripts/README.md
Normal file
@@ -0,0 +1,76 @@
|
|||||||
|
# scripts 工具脚本说明
|
||||||
|
|
||||||
|
本目录包含数据生成、回放、可视化与分析等工具脚本。训练脚本(`train_bc.py`、`train_magail.py`)位于项目根目录。
|
||||||
|
|
||||||
|
## 路径约定(相对项目根)
|
||||||
|
|
||||||
|
- **数据**:`data/exp_filtered`(Waymo 场景)、`data/training_data`(专家 pkl 输出)
|
||||||
|
- **模型**:`models/bc/`(BC)、`models/magail/`(MAGAIL)
|
||||||
|
- **日志**:`logs/bc/`、`logs/magail/`(TensorBoard)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 脚本列表与用法
|
||||||
|
|
||||||
|
### 数据生成
|
||||||
|
|
||||||
|
| 脚本 | 用途 | 用法示例 |
|
||||||
|
|------|------|----------|
|
||||||
|
| [generate_expert_data.py](generate_expert_data.py) | 从 Waymo 数据生成专家 (obs, act) 的 pkl | `python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100` |
|
||||||
|
|
||||||
|
**常用参数**:`--data_dir`(默认 `data/exp_filtered`)、`--output_dir`(默认 `data/training_data`)、`--start_index`、`--num_scenarios`。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 可视化(统一入口)
|
||||||
|
|
||||||
|
| 脚本 | 用途 | 用法示例 |
|
||||||
|
|------|------|----------|
|
||||||
|
| [visualize.py](visualize.py) | **replay**:场景回放(ExpertReplayEnv);**policy**:BC/MAGAIL 策略;**trajectory**:专家轨迹 2D 动画 | 见下方 |
|
||||||
|
|
||||||
|
**子命令**:
|
||||||
|
|
||||||
|
- **replay**(原始专家轨迹回放):
|
||||||
|
```bash
|
||||||
|
python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500
|
||||||
|
```
|
||||||
|
|
||||||
|
- **policy**(BC 或 MAGAIL 训练策略):
|
||||||
|
```bash
|
||||||
|
python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1
|
||||||
|
python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth --num_scenarios 1 --deterministic
|
||||||
|
```
|
||||||
|
|
||||||
|
- **trajectory**(专家轨迹 matplotlib 俯视图动画):
|
||||||
|
```bash
|
||||||
|
python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_idx 0
|
||||||
|
```
|
||||||
|
|
||||||
|
**公共参数**:`--data_dir`(默认 `data/exp_filtered`)、`--start_index`、`--num_scenarios`、`--horizon`。policy 模式另有 `--policy_type`(auto/bc/magail)、`--model_path`、`--deterministic`(仅 MAGAIL)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 数据分析与检查
|
||||||
|
|
||||||
|
| 脚本 | 用途 | 用法示例 |
|
||||||
|
|------|------|----------|
|
||||||
|
| [analyze_expert_data.py](analyze_expert_data.py) | 分析专家数据分布与统计 | 见脚本内 `__main__`(依赖 env 与数据目录配置) |
|
||||||
|
| [check_track_fields.py](check_track_fields.py) | 检查 Waymo 轨迹字段 | 见脚本内 `__main__` |
|
||||||
|
| [check_database_info.py](check_database_info.py) | 检查数据库/场景信息 | 见脚本内 `__main__`(含硬编码路径,可按需改为 `data/exp_filtered`) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 其他
|
||||||
|
|
||||||
|
| 脚本 | 用途 | 用法示例 |
|
||||||
|
|------|------|----------|
|
||||||
|
| [launch_tensorboard.py](launch_tensorboard.py) | 启动 TensorBoard | `python scripts/launch_tensorboard.py --logdir logs`(或 `logs/bc` / `logs/magail`) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 与训练流程的对应关系
|
||||||
|
|
||||||
|
1. **数据准备**:`generate_expert_data.py` → 输出到 `data/training_data/*.pkl`
|
||||||
|
2. **BC 训练**:根目录 `train_bc.py` → 模型保存到 `models/bc/`,日志到 `logs/bc/`
|
||||||
|
3. **MAGAIL 训练**:根目录 `train_magail.py` → 模型保存到 `models/magail/`,日志到 `logs/magail/`
|
||||||
|
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
||||||
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()
|
||||||
162
scripts/generate_expert_data.py
Normal file
162
scripts/generate_expert_data.py
Normal file
@@ -0,0 +1,162 @@
|
|||||||
|
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 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
|
||||||
|
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)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
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())
|
||||||
395
scripts/visualize.py
Normal file
395
scripts/visualize.py
Normal file
@@ -0,0 +1,395 @@
|
|||||||
|
"""
|
||||||
|
Unified visualization: replay (scenario replay), policy (BC/MAGAIL), trajectory (2D expert trajectory animation).
|
||||||
|
Usage: python scripts/visualize.py <replay|policy|trajectory> [args...]
|
||||||
|
"""
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
if project_root not in sys.path:
|
||||||
|
sys.path.insert(0, project_root)
|
||||||
|
|
||||||
|
# --- Replay ---
|
||||||
|
def _run_replay(args):
|
||||||
|
from Env.expert_replay_env import ExpertReplayEnv
|
||||||
|
|
||||||
|
data_path = os.path.abspath(args.data_dir)
|
||||||
|
if not os.path.exists(data_path):
|
||||||
|
raise ValueError(f"Data directory {data_path} not found")
|
||||||
|
|
||||||
|
from metadrive.scenario.utils import read_dataset_summary
|
||||||
|
_, summary_lookup, _ = read_dataset_summary(data_path)
|
||||||
|
if args.start_index >= len(summary_lookup):
|
||||||
|
raise ValueError(
|
||||||
|
f"start_index={args.start_index} out of range. Dataset has {len(summary_lookup)} scenarios."
|
||||||
|
)
|
||||||
|
max_available = len(summary_lookup) - args.start_index
|
||||||
|
num_to_run = min(args.num_scenarios, max_available)
|
||||||
|
|
||||||
|
env_config = {
|
||||||
|
"data_directory": data_path,
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"num_controlled_agents": 100,
|
||||||
|
"horizon": args.horizon,
|
||||||
|
"use_render": True,
|
||||||
|
"sequential_seed": True,
|
||||||
|
"reactive_traffic": False,
|
||||||
|
"start_scenario_index": args.start_index,
|
||||||
|
"num_scenarios": -1,
|
||||||
|
"log_level": 40,
|
||||||
|
}
|
||||||
|
|
||||||
|
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||||
|
env = ExpertReplayEnv(config=env_config)
|
||||||
|
|
||||||
|
try:
|
||||||
|
for i in range(args.start_index, args.start_index + num_to_run):
|
||||||
|
print(f"\n--- Playing Scenario {i} ---")
|
||||||
|
try:
|
||||||
|
obs = env.reset(seed=i)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error resetting scenario {i}: {e}")
|
||||||
|
continue
|
||||||
|
|
||||||
|
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
||||||
|
|
||||||
|
for step in range(args.horizon):
|
||||||
|
obs, rewards, dones, infos = env.step(None)
|
||||||
|
env.render(
|
||||||
|
mode="top_down",
|
||||||
|
text={"Step": step, "Agents": len(env.controlled_agents), "Scenario": i},
|
||||||
|
)
|
||||||
|
time.sleep(0.05)
|
||||||
|
if dones["__all__"]:
|
||||||
|
print(f"Scenario {i} finished at step {step}")
|
||||||
|
break
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("Interrupted by user")
|
||||||
|
except Exception as e:
|
||||||
|
import traceback
|
||||||
|
traceback.print_exc()
|
||||||
|
print(f"Global error: {e}")
|
||||||
|
finally:
|
||||||
|
env.close()
|
||||||
|
print("Environment closed.")
|
||||||
|
|
||||||
|
|
||||||
|
# --- Policy (BC / MAGAIL) ---
|
||||||
|
def _resolve_data_dir(data_dir_arg):
|
||||||
|
if data_dir_arg:
|
||||||
|
data_dir = data_dir_arg
|
||||||
|
else:
|
||||||
|
data_dir = os.path.join(project_root, "data", "exp_filtered")
|
||||||
|
if not os.path.exists(data_dir):
|
||||||
|
data_dir = os.path.join(project_root, "data", "exp_converted")
|
||||||
|
if not os.path.exists(data_dir):
|
||||||
|
raise FileNotFoundError(f"Data directory not found at {data_dir}. Please specify --data_dir.")
|
||||||
|
return data_dir
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_model_path(model_path, policy_type):
|
||||||
|
if os.path.exists(model_path):
|
||||||
|
return model_path
|
||||||
|
if policy_type == "bc":
|
||||||
|
candidate = os.path.join(project_root, "models", "bc", os.path.basename(model_path))
|
||||||
|
else:
|
||||||
|
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path))
|
||||||
|
if os.path.exists(candidate):
|
||||||
|
return candidate
|
||||||
|
if policy_type == "magail" and not model_path.endswith("_actor.pth"):
|
||||||
|
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path) + "_actor.pth")
|
||||||
|
if os.path.exists(candidate):
|
||||||
|
return candidate
|
||||||
|
raise FileNotFoundError(f"Model path {model_path} not found.")
|
||||||
|
|
||||||
|
|
||||||
|
def _run_policy(args):
|
||||||
|
from Env.bc_env import BCScenarioEnv
|
||||||
|
from metadrive.engine.engine_utils import close_engine
|
||||||
|
|
||||||
|
policy_type = (args.policy_type or "auto").lower()
|
||||||
|
if policy_type == "auto":
|
||||||
|
policy_type = "bc" if args.model_path.endswith(".pt") else "magail"
|
||||||
|
|
||||||
|
data_dir = _resolve_data_dir(args.data_dir)
|
||||||
|
data_path = os.path.abspath(data_dir)
|
||||||
|
env_config = {
|
||||||
|
"data_directory": data_path,
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"num_controlled_agents": 3,
|
||||||
|
"horizon": args.horizon,
|
||||||
|
"use_render": True,
|
||||||
|
"sequential_seed": True,
|
||||||
|
"start_scenario_index": args.start_index,
|
||||||
|
"num_scenarios": args.num_scenarios,
|
||||||
|
"log_level": 40,
|
||||||
|
}
|
||||||
|
|
||||||
|
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
||||||
|
try:
|
||||||
|
env = BCScenarioEnv(env_config, agent2policy={})
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error init env: {e}. Trying to close lingering engine...")
|
||||||
|
try:
|
||||||
|
close_engine()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
env = BCScenarioEnv(env_config, agent2policy={})
|
||||||
|
|
||||||
|
state_dim = 45
|
||||||
|
action_dim = 2
|
||||||
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
model_path = _resolve_model_path(args.model_path, policy_type)
|
||||||
|
print(f"Loading model from {model_path}...")
|
||||||
|
|
||||||
|
if policy_type == "bc":
|
||||||
|
from Algorithm.policy import StateIndependentPolicy
|
||||||
|
policy = StateIndependentPolicy(
|
||||||
|
state_shape=(state_dim,),
|
||||||
|
action_shape=(action_dim,),
|
||||||
|
hidden_units=(256, 256),
|
||||||
|
hidden_activation=torch.nn.Tanh(),
|
||||||
|
).to(device)
|
||||||
|
policy.load_state_dict(torch.load(model_path, map_location=device))
|
||||||
|
policy.eval()
|
||||||
|
else:
|
||||||
|
from train_magail import Actor
|
||||||
|
actor = Actor(state_dim, action_dim).to(device)
|
||||||
|
actor.load_state_dict(torch.load(model_path, map_location=device))
|
||||||
|
actor.eval()
|
||||||
|
|
||||||
|
try:
|
||||||
|
for i in range(args.start_index, args.start_index + args.num_scenarios):
|
||||||
|
print(f"\n--- Playing Scenario {i} ---")
|
||||||
|
try:
|
||||||
|
obs_dict = env.reset(seed=i)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error resetting {i}: {e}. Skipping.")
|
||||||
|
try:
|
||||||
|
close_engine()
|
||||||
|
env = BCScenarioEnv(env_config, agent2policy={})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
continue
|
||||||
|
|
||||||
|
print(f"Scenario loaded. Controlled agents: {len(obs_dict)}")
|
||||||
|
step_count = 0
|
||||||
|
episode_reward = 0.0
|
||||||
|
|
||||||
|
while True:
|
||||||
|
agent_ids = list(obs_dict.keys())
|
||||||
|
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||||
|
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
if policy_type == "bc":
|
||||||
|
actions_np = policy(obs_tensor).cpu().numpy()
|
||||||
|
else:
|
||||||
|
dist = actor(obs_tensor)
|
||||||
|
if args.deterministic:
|
||||||
|
actions_np = torch.tanh(dist.mean).cpu().numpy()
|
||||||
|
else:
|
||||||
|
actions_np = torch.tanh(dist.sample()).cpu().numpy()
|
||||||
|
|
||||||
|
actions = {aid: actions_np[idx].flatten() for idx, aid in enumerate(agent_ids)}
|
||||||
|
obs_dict, rewards, dones, infos = env.step(actions)
|
||||||
|
episode_reward += sum(rewards.values())
|
||||||
|
|
||||||
|
env.render(
|
||||||
|
mode="top_down",
|
||||||
|
text={
|
||||||
|
"Scenario": i,
|
||||||
|
"Step": step_count,
|
||||||
|
"Agents": len(obs_dict),
|
||||||
|
"Total Reward": f"{episode_reward:.2f}",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
step_count += 1
|
||||||
|
|
||||||
|
if dones["__all__"] or step_count >= args.horizon:
|
||||||
|
print(f"Scenario finished at step {step_count}, reward {episode_reward:.2f}")
|
||||||
|
break
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("Interrupted.")
|
||||||
|
finally:
|
||||||
|
env.close()
|
||||||
|
|
||||||
|
|
||||||
|
# --- Trajectory (matplotlib 2D animation) ---
|
||||||
|
def _build_expert_trajectories_from_env(env):
|
||||||
|
"""Build expert_trajectories dict from env (ExpertReplayEnv has traffic_manager.current_traffic_data)."""
|
||||||
|
if hasattr(env, "expert_trajectories") and env.expert_trajectories:
|
||||||
|
return env.expert_trajectories
|
||||||
|
if not hasattr(env, "engine") or not hasattr(env.engine, "traffic_manager"):
|
||||||
|
return {}
|
||||||
|
from metadrive.type import MetaDriveType
|
||||||
|
data = getattr(env.engine.traffic_manager, "current_traffic_data", None)
|
||||||
|
if not data:
|
||||||
|
return {}
|
||||||
|
expert_trajs = {}
|
||||||
|
for scenario_id, track in data.items():
|
||||||
|
if track.get("type") != MetaDriveType.VEHICLE or "state" not in track:
|
||||||
|
continue
|
||||||
|
state = track["state"]
|
||||||
|
positions = state.get("position")
|
||||||
|
if positions is None:
|
||||||
|
continue
|
||||||
|
valid = state.get("valid", np.ones(len(positions), dtype=bool))
|
||||||
|
valid = np.asarray(valid).flatten()
|
||||||
|
if valid.size != len(positions):
|
||||||
|
valid = np.ones(len(positions), dtype=bool)
|
||||||
|
first_show = int(np.argmax(valid)) if valid.any() else 0
|
||||||
|
last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else len(positions) - 1
|
||||||
|
obj_id = track.get("metadata", {}).get("object_id", str(scenario_id))
|
||||||
|
expert_trajs[obj_id] = {
|
||||||
|
"positions": np.asarray(positions),
|
||||||
|
"start_timestep": first_show,
|
||||||
|
"end_timestep": last_show,
|
||||||
|
}
|
||||||
|
return expert_trajs
|
||||||
|
|
||||||
|
|
||||||
|
def _run_trajectory_animation(expert_trajs, scenario_idx):
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
from matplotlib.animation import FuncAnimation
|
||||||
|
|
||||||
|
if len(expert_trajs) == 0:
|
||||||
|
print("No expert trajectories to visualize.")
|
||||||
|
return
|
||||||
|
|
||||||
|
fig, ax = plt.subplots(figsize=(12, 12))
|
||||||
|
max_timestep = max(t["end_timestep"] for t in expert_trajs.values())
|
||||||
|
min_timestep = min(t["start_timestep"] for t in expert_trajs.values())
|
||||||
|
|
||||||
|
colors = plt.cm.tab10(np.linspace(0, 1, len(expert_trajs)))
|
||||||
|
for idx, (obj_id, traj) in enumerate(expert_trajs.items()):
|
||||||
|
positions = np.asarray(traj["positions"])
|
||||||
|
if positions.ndim >= 2:
|
||||||
|
positions = positions[:, :2]
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
ax.plot(
|
||||||
|
positions[:, 0], positions[:, 1],
|
||||||
|
color=colors[idx], alpha=0.3, linewidth=1,
|
||||||
|
label=f"Vehicle {str(obj_id)[:6]}",
|
||||||
|
)
|
||||||
|
|
||||||
|
scatter = ax.scatter([], [], s=200, c="red", marker="o", edgecolors="black", linewidths=2)
|
||||||
|
time_text = ax.text(0.02, 0.95, "", transform=ax.transAxes, fontsize=14)
|
||||||
|
ax.set_xlabel("X (m)")
|
||||||
|
ax.set_ylabel("Y (m)")
|
||||||
|
ax.set_title(f"Expert Trajectory Visualization - Scenario {scenario_idx}")
|
||||||
|
ax.legend(loc="upper right", fontsize=8)
|
||||||
|
ax.grid(True, alpha=0.3)
|
||||||
|
ax.axis("equal")
|
||||||
|
|
||||||
|
def update(frame):
|
||||||
|
current_time = min_timestep + frame
|
||||||
|
current_positions = []
|
||||||
|
for traj in expert_trajs.values():
|
||||||
|
st, et = traj["start_timestep"], traj["end_timestep"]
|
||||||
|
if st <= current_time <= et:
|
||||||
|
pos = np.asarray(traj["positions"])
|
||||||
|
if pos.ndim >= 2:
|
||||||
|
pos = pos[current_time - st, :2]
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
current_positions.append(pos)
|
||||||
|
if current_positions:
|
||||||
|
scatter.set_offsets(np.array(current_positions))
|
||||||
|
time_text.set_text(f"Time: {frame * 0.1:.1f}s (Frame {frame})")
|
||||||
|
return scatter, time_text
|
||||||
|
|
||||||
|
anim = FuncAnimation(
|
||||||
|
fig, update, frames=max_timestep - min_timestep + 1,
|
||||||
|
interval=100, blit=True, repeat=True,
|
||||||
|
)
|
||||||
|
plt.tight_layout()
|
||||||
|
plt.show()
|
||||||
|
return anim
|
||||||
|
|
||||||
|
|
||||||
|
def _run_trajectory(args):
|
||||||
|
from Env.expert_replay_env import ExpertReplayEnv
|
||||||
|
|
||||||
|
data_dir = _resolve_data_dir(args.data_dir)
|
||||||
|
data_path = os.path.abspath(data_dir)
|
||||||
|
env_config = {
|
||||||
|
"data_directory": data_path,
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"num_controlled_agents": 100,
|
||||||
|
"horizon": 500,
|
||||||
|
"use_render": False,
|
||||||
|
"sequential_seed": True,
|
||||||
|
"reactive_traffic": False,
|
||||||
|
"start_scenario_index": args.scenario_idx,
|
||||||
|
"num_scenarios": 1,
|
||||||
|
"log_level": 40,
|
||||||
|
}
|
||||||
|
|
||||||
|
env = ExpertReplayEnv(config=env_config)
|
||||||
|
try:
|
||||||
|
env.reset(seed=args.scenario_idx)
|
||||||
|
expert_trajs = _build_expert_trajectories_from_env(env)
|
||||||
|
_run_trajectory_animation(expert_trajs, args.scenario_idx)
|
||||||
|
finally:
|
||||||
|
env.close()
|
||||||
|
|
||||||
|
|
||||||
|
# --- Main ---
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Unified visualization: replay, policy (BC/MAGAIL), trajectory.",
|
||||||
|
)
|
||||||
|
subparsers = parser.add_subparsers(dest="mode", required=True, help="replay | policy | trajectory")
|
||||||
|
|
||||||
|
# Common args for data_dir (used by all)
|
||||||
|
def add_common_data_args(p):
|
||||||
|
p.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Waymo scenario directory")
|
||||||
|
p.add_argument("--start_index", type=int, default=0)
|
||||||
|
p.add_argument("--num_scenarios", type=int, default=1)
|
||||||
|
p.add_argument("--horizon", type=int, default=200)
|
||||||
|
|
||||||
|
# replay
|
||||||
|
pr = subparsers.add_parser("replay", help="Replay scenario with ExpertReplayEnv (no policy)")
|
||||||
|
add_common_data_args(pr)
|
||||||
|
pr.set_defaults(horizon=500)
|
||||||
|
|
||||||
|
# policy
|
||||||
|
pp = subparsers.add_parser("policy", help="Visualize BC or MAGAIL trained policy")
|
||||||
|
add_common_data_args(pp)
|
||||||
|
pp.add_argument("--policy_type", type=str, default="auto", choices=["auto", "bc", "magail"])
|
||||||
|
pp.add_argument("--model_path", type=str, default="models/bc/policy_best.pt")
|
||||||
|
pp.add_argument("--deterministic", action="store_true", help="MAGAIL: use mean action")
|
||||||
|
|
||||||
|
# trajectory
|
||||||
|
pt = subparsers.add_parser("trajectory", help="2D matplotlib animation of expert trajectories")
|
||||||
|
pt.add_argument("--data_dir", type=str, default="data/exp_filtered")
|
||||||
|
pt.add_argument("--scenario_idx", type=int, default=0)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Resolve data_dir relative to project root when default
|
||||||
|
if args.mode != "trajectory":
|
||||||
|
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||||
|
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||||
|
else:
|
||||||
|
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||||
|
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||||
|
|
||||||
|
if args.mode == "replay":
|
||||||
|
_run_replay(args)
|
||||||
|
elif args.mode == "policy":
|
||||||
|
_run_policy(args)
|
||||||
|
elif args.mode == "trajectory":
|
||||||
|
_run_trajectory(args)
|
||||||
|
else:
|
||||||
|
parser.error(f"Unknown mode: {args.mode}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
146
train_bc.py
Normal file
146
train_bc.py
Normal file
@@ -0,0 +1,146 @@
|
|||||||
|
"""
|
||||||
|
BC 训练脚本:负责数据加载、环境评估、日志与保存;BC 算法由 Algorithm.bc 提供。
|
||||||
|
使用方式不变:python train_bc.py [--expert_data_path data/training_data] [--save_dir models/bc] ...
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import argparse
|
||||||
|
from torch.utils.data import DataLoader, TensorDataset
|
||||||
|
from torch.optim import Adam
|
||||||
|
from torch.optim.lr_scheduler import ExponentialLR
|
||||||
|
from datetime import datetime
|
||||||
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
|
|
||||||
|
from Algorithm.policy import StateIndependentPolicy
|
||||||
|
from Algorithm.bc import train_bc_epoch, eval_bc_epoch
|
||||||
|
from Env.bc_env import BCScenarioEnv
|
||||||
|
from dataset.loader import load_expert_pkl
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate_policy(policy, args, device):
|
||||||
|
"""在 BCScenarioEnv 中评估策略,跑若干 episode,返回平均 reward。"""
|
||||||
|
waymo_data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
|
||||||
|
data_dir = os.path.join(waymo_data_dir, "exp_filtered")
|
||||||
|
if not os.path.exists(data_dir):
|
||||||
|
data_dir = os.path.join(waymo_data_dir, "exp_converted")
|
||||||
|
if not os.path.exists(data_dir):
|
||||||
|
print(f"[ERROR] Could not find scenario data in {waymo_data_dir}. Evaluation skipped.")
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
env_config = {
|
||||||
|
"data_directory": data_dir,
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"num_controlled_agents": 3,
|
||||||
|
"use_render": False,
|
||||||
|
"sequential_seed": True,
|
||||||
|
"horizon": 200,
|
||||||
|
}
|
||||||
|
env = BCScenarioEnv(env_config, agent2policy=None)
|
||||||
|
total_rewards = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
for i in range(3):
|
||||||
|
obs_dict = env.reset(seed=i)
|
||||||
|
episode_reward = 0
|
||||||
|
dones = {"__all__": False}
|
||||||
|
step_count = 0
|
||||||
|
horizon = 200
|
||||||
|
while not dones["__all__"]:
|
||||||
|
step_count += 1
|
||||||
|
if step_count >= horizon:
|
||||||
|
break
|
||||||
|
if not obs_dict:
|
||||||
|
obs_dict, _, dones, _ = env.step({})
|
||||||
|
continue
|
||||||
|
agent_ids = list(obs_dict.keys())
|
||||||
|
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||||
|
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||||
|
with torch.no_grad():
|
||||||
|
actions, _ = policy.sample(obs_tensor)
|
||||||
|
actions = actions.cpu().numpy()
|
||||||
|
action_dict = {aid: act for aid, act in zip(agent_ids, actions)}
|
||||||
|
obs_dict, rewards, dones, _ = env.step(action_dict)
|
||||||
|
episode_reward += sum(rewards.values())
|
||||||
|
total_rewards.append(episode_reward)
|
||||||
|
print(f" Eval Episode {i}: Total Reward {episode_reward:.2f}")
|
||||||
|
avg_reward = float(np.mean(total_rewards))
|
||||||
|
print(f" Average Evaluation Reward: {avg_reward:.2f}")
|
||||||
|
return avg_reward
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Evaluation failed: {e}")
|
||||||
|
import traceback
|
||||||
|
traceback.print_exc()
|
||||||
|
return 0.0
|
||||||
|
finally:
|
||||||
|
env.close()
|
||||||
|
|
||||||
|
|
||||||
|
def main(args):
|
||||||
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
print(f"Using device: {device}")
|
||||||
|
|
||||||
|
os.makedirs("logs/bc", exist_ok=True)
|
||||||
|
log_dir = os.path.join("logs", "bc", datetime.now().strftime("%Y%m%d-%H%M%S"))
|
||||||
|
writer = SummaryWriter(log_dir)
|
||||||
|
print(f"TensorBoard logging to: {log_dir}")
|
||||||
|
os.makedirs(args.save_dir, exist_ok=True)
|
||||||
|
|
||||||
|
obs_data, act_data = load_expert_pkl(args.expert_data_path)
|
||||||
|
obs_tensor = torch.FloatTensor(obs_data)
|
||||||
|
act_tensor = torch.FloatTensor(act_data)
|
||||||
|
dataset = TensorDataset(obs_tensor, act_tensor)
|
||||||
|
train_size = int(0.8 * len(dataset))
|
||||||
|
val_size = len(dataset) - train_size
|
||||||
|
train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])
|
||||||
|
train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True)
|
||||||
|
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False)
|
||||||
|
print(f"Dataset loaded. Train size: {len(train_dataset)}, Val size: {len(val_dataset)}")
|
||||||
|
|
||||||
|
state_dim = obs_data.shape[1]
|
||||||
|
action_dim = act_data.shape[1]
|
||||||
|
print(f"State Dim: {state_dim}, Action Dim: {action_dim}")
|
||||||
|
|
||||||
|
policy = StateIndependentPolicy(
|
||||||
|
state_shape=(state_dim,),
|
||||||
|
action_shape=(action_dim,),
|
||||||
|
hidden_units=(256, 256),
|
||||||
|
hidden_activation=torch.nn.Tanh(),
|
||||||
|
).to(device)
|
||||||
|
optimizer = Adam(policy.parameters(), lr=args.lr)
|
||||||
|
scheduler = ExponentialLR(optimizer, gamma=0.99)
|
||||||
|
|
||||||
|
best_val_loss = float("inf")
|
||||||
|
for epoch in range(args.epochs):
|
||||||
|
avg_train_loss = train_bc_epoch(policy, train_loader, optimizer, device)
|
||||||
|
scheduler.step()
|
||||||
|
avg_val_loss = eval_bc_epoch(policy, val_loader, device)
|
||||||
|
|
||||||
|
print(f"Epoch {epoch+1}/{args.epochs} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}")
|
||||||
|
writer.add_scalar("Loss/train", avg_train_loss, epoch)
|
||||||
|
writer.add_scalar("Loss/val", avg_val_loss, epoch)
|
||||||
|
writer.add_scalar("Learning_rate", scheduler.get_last_lr()[0], epoch)
|
||||||
|
|
||||||
|
if avg_val_loss < best_val_loss:
|
||||||
|
best_val_loss = avg_val_loss
|
||||||
|
torch.save(policy.state_dict(), os.path.join(args.save_dir, "policy_best.pt"))
|
||||||
|
|
||||||
|
if (epoch + 1) % args.eval_freq == 0:
|
||||||
|
eval_reward = evaluate_policy(policy, args, device)
|
||||||
|
writer.add_scalar("Reward/eval", eval_reward, epoch)
|
||||||
|
|
||||||
|
torch.save(policy.state_dict(), os.path.join(args.save_dir, "policy_final.pt"))
|
||||||
|
writer.close()
|
||||||
|
print("Training finished.")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--expert_data_path", type=str, default="data/training_data", help="Path to expert data pickle or directory")
|
||||||
|
parser.add_argument("--save_dir", type=str, default="models/bc", help="Directory to save models")
|
||||||
|
parser.add_argument("--epochs", type=int, default=100)
|
||||||
|
parser.add_argument("--batch_size", type=int, default=64)
|
||||||
|
parser.add_argument("--lr", type=float, default=3e-4)
|
||||||
|
parser.add_argument("--eval_freq", type=int, default=10)
|
||||||
|
args = parser.parse_args()
|
||||||
|
main(args)
|
||||||
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