完善项目目录结构
This commit is contained in:
@@ -22,36 +22,31 @@
|
||||
|
||||
---
|
||||
|
||||
### 回放与可视化
|
||||
### 可视化(统一入口)
|
||||
|
||||
| 脚本 | 用途 | 用法示例 |
|
||||
|------|------|----------|
|
||||
| [visualize_replay.py](visualize_replay.py) | 原始专家轨迹回放(ExpertReplayEnv) | `python scripts/visualize_replay.py --data_dir data/exp_filtered --num_scenarios 1 --horizon 200` |
|
||||
| [visualize_trained_policy.py](visualize_trained_policy.py) | **BC/MAGAIL 共用**:加载训练好的策略在 45 维场景中可视化 | 见下方「训练策略可视化」小节 |
|
||||
| [visualize.py](visualize.py) | **replay**:场景回放(ExpertReplayEnv);**policy**:BC/MAGAIL 策略;**trajectory**:专家轨迹 2D 动画 | 见下方 |
|
||||
|
||||
#### 训练策略可视化(visualize_trained_policy.py)
|
||||
**子命令**:
|
||||
|
||||
使用训练好的 **BC** 或 **MAGAIL** 模型在 45 维场景环境中运行,并实时渲染俯瞰图(top-down view)。统一入口:`scripts/visualize_trained_policy.py`。
|
||||
|
||||
**BC 模型**:
|
||||
- **replay**(原始专家轨迹回放):
|
||||
```bash
|
||||
python scripts/visualize_trained_policy.py --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1
|
||||
python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500
|
||||
```
|
||||
|
||||
**MAGAIL 模型**:
|
||||
- **policy**(BC 或 MAGAIL 训练策略):
|
||||
```bash
|
||||
python scripts/visualize_trained_policy.py --policy_type magail --model_path models/magail/model_50_actor.pth --data_dir data/exp_filtered --num_scenarios 1 --deterministic
|
||||
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
|
||||
```
|
||||
|
||||
**自动推断类型**(根据 `--model_path` 扩展名:`.pt` → BC,否则 → MAGAIL):
|
||||
- **trajectory**(专家轨迹 matplotlib 俯视图动画):
|
||||
```bash
|
||||
python scripts/visualize_trained_policy.py --model_path models/bc/policy_best.pt
|
||||
python scripts/visualize_trained_policy.py --model_path models/magail/model_50_actor.pth
|
||||
python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_idx 0
|
||||
```
|
||||
|
||||
**根目录 BC 薄包装**:`python visualize_bc.py --model_path models/bc/policy_best.pt`
|
||||
|
||||
**参数**:`--policy_type`(`auto`|`bc`|`magail`)、`--model_path`(默认 `models/bc/policy_best.pt`)、`--data_dir`、`--start_index`、`--num_scenarios`、`--horizon`、`--deterministic`(仅 MAGAIL)。环境统一为 45 维 `BCScenarioEnv`,渲染为 MetaDrive top_down。数据目录未指定时默认 `data/exp_filtered`(不存在则 `data/exp_converted`)。
|
||||
**公共参数**:`--data_dir`(默认 `data/exp_filtered`)、`--start_index`、`--num_scenarios`、`--horizon`。policy 模式另有 `--policy_type`(auto/bc/magail)、`--model_path`、`--deterministic`(仅 MAGAIL)。
|
||||
|
||||
---
|
||||
|
||||
@@ -62,7 +57,6 @@ python scripts/visualize_trained_policy.py --model_path models/magail/model_50_a
|
||||
| [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`) |
|
||||
| [visualize_expert_trajectory.py](visualize_expert_trajectory.py) | 用 matplotlib 画专家轨迹动画 | 依赖 `env.expert_trajectories`,与当前 env 接口可能不一致,可选使用 |
|
||||
|
||||
---
|
||||
|
||||
@@ -79,4 +73,4 @@ python scripts/visualize_trained_policy.py --model_path models/magail/model_50_a
|
||||
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. **可视化**:`visualize_trained_policy.py`(或根目录 `visualize_bc.py` 仅 BC)→ 从 `models/bc` 或 `models/magail` 加载模型,数据目录默认 `data/exp_filtered`
|
||||
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
||||
|
||||
395
scripts/visualize.py
Normal file
395
scripts/visualize.py
Normal file
@@ -0,0 +1,395 @@
|
||||
"""
|
||||
Unified visualization: replay (scenario replay), policy (BC/MAGAIL), trajectory (2D expert trajectory animation).
|
||||
Usage: python scripts/visualize.py <replay|policy|trajectory> [args...]
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
# --- Replay ---
|
||||
def _run_replay(args):
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
data_path = os.path.abspath(args.data_dir)
|
||||
if not os.path.exists(data_path):
|
||||
raise ValueError(f"Data directory {data_path} not found")
|
||||
|
||||
from metadrive.scenario.utils import read_dataset_summary
|
||||
_, summary_lookup, _ = read_dataset_summary(data_path)
|
||||
if args.start_index >= len(summary_lookup):
|
||||
raise ValueError(
|
||||
f"start_index={args.start_index} out of range. Dataset has {len(summary_lookup)} scenarios."
|
||||
)
|
||||
max_available = len(summary_lookup) - args.start_index
|
||||
num_to_run = min(args.num_scenarios, max_available)
|
||||
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": -1,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + num_to_run):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
||||
|
||||
for step in range(args.horizon):
|
||||
obs, rewards, dones, infos = env.step(None)
|
||||
env.render(
|
||||
mode="top_down",
|
||||
text={"Step": step, "Agents": len(env.controlled_agents), "Scenario": i},
|
||||
)
|
||||
time.sleep(0.05)
|
||||
if dones["__all__"]:
|
||||
print(f"Scenario {i} finished at step {step}")
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted by user")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
env.close()
|
||||
print("Environment closed.")
|
||||
|
||||
|
||||
# --- Policy (BC / MAGAIL) ---
|
||||
def _resolve_data_dir(data_dir_arg):
|
||||
if data_dir_arg:
|
||||
data_dir = data_dir_arg
|
||||
else:
|
||||
data_dir = os.path.join(project_root, "data", "exp_filtered")
|
||||
if not os.path.exists(data_dir):
|
||||
data_dir = os.path.join(project_root, "data", "exp_converted")
|
||||
if not os.path.exists(data_dir):
|
||||
raise FileNotFoundError(f"Data directory not found at {data_dir}. Please specify --data_dir.")
|
||||
return data_dir
|
||||
|
||||
|
||||
def _resolve_model_path(model_path, policy_type):
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
if policy_type == "bc":
|
||||
candidate = os.path.join(project_root, "models", "bc", os.path.basename(model_path))
|
||||
else:
|
||||
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path))
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
if policy_type == "magail" and not model_path.endswith("_actor.pth"):
|
||||
candidate = os.path.join(project_root, "models", "magail", os.path.basename(model_path) + "_actor.pth")
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
raise FileNotFoundError(f"Model path {model_path} not found.")
|
||||
|
||||
|
||||
def _run_policy(args):
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
from metadrive.engine.engine_utils import close_engine
|
||||
|
||||
policy_type = (args.policy_type or "auto").lower()
|
||||
if policy_type == "auto":
|
||||
policy_type = "bc" if args.model_path.endswith(".pt") else "magail"
|
||||
|
||||
data_dir = _resolve_data_dir(args.data_dir)
|
||||
data_path = os.path.abspath(data_dir)
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True,
|
||||
"sequential_seed": True,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": args.num_scenarios,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
||||
try:
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
except Exception as e:
|
||||
print(f"Error init env: {e}. Trying to close lingering engine...")
|
||||
try:
|
||||
close_engine()
|
||||
except Exception:
|
||||
pass
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
|
||||
state_dim = 45
|
||||
action_dim = 2
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
model_path = _resolve_model_path(args.model_path, policy_type)
|
||||
print(f"Loading model from {model_path}...")
|
||||
|
||||
if policy_type == "bc":
|
||||
from Algorithm.policy import StateIndependentPolicy
|
||||
policy = StateIndependentPolicy(
|
||||
state_shape=(state_dim,),
|
||||
action_shape=(action_dim,),
|
||||
hidden_units=(256, 256),
|
||||
hidden_activation=torch.nn.Tanh(),
|
||||
).to(device)
|
||||
policy.load_state_dict(torch.load(model_path, map_location=device))
|
||||
policy.eval()
|
||||
else:
|
||||
from train_magail import Actor
|
||||
actor = Actor(state_dim, action_dim).to(device)
|
||||
actor.load_state_dict(torch.load(model_path, map_location=device))
|
||||
actor.eval()
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + args.num_scenarios):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs_dict = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting {i}: {e}. Skipping.")
|
||||
try:
|
||||
close_engine()
|
||||
env = BCScenarioEnv(env_config, agent2policy={})
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(obs_dict)}")
|
||||
step_count = 0
|
||||
episode_reward = 0.0
|
||||
|
||||
while True:
|
||||
agent_ids = list(obs_dict.keys())
|
||||
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
if policy_type == "bc":
|
||||
actions_np = policy(obs_tensor).cpu().numpy()
|
||||
else:
|
||||
dist = actor(obs_tensor)
|
||||
if args.deterministic:
|
||||
actions_np = torch.tanh(dist.mean).cpu().numpy()
|
||||
else:
|
||||
actions_np = torch.tanh(dist.sample()).cpu().numpy()
|
||||
|
||||
actions = {aid: actions_np[idx].flatten() for idx, aid in enumerate(agent_ids)}
|
||||
obs_dict, rewards, dones, infos = env.step(actions)
|
||||
episode_reward += sum(rewards.values())
|
||||
|
||||
env.render(
|
||||
mode="top_down",
|
||||
text={
|
||||
"Scenario": i,
|
||||
"Step": step_count,
|
||||
"Agents": len(obs_dict),
|
||||
"Total Reward": f"{episode_reward:.2f}",
|
||||
},
|
||||
)
|
||||
step_count += 1
|
||||
|
||||
if dones["__all__"] or step_count >= args.horizon:
|
||||
print(f"Scenario finished at step {step_count}, reward {episode_reward:.2f}")
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted.")
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
# --- Trajectory (matplotlib 2D animation) ---
|
||||
def _build_expert_trajectories_from_env(env):
|
||||
"""Build expert_trajectories dict from env (ExpertReplayEnv has traffic_manager.current_traffic_data)."""
|
||||
if hasattr(env, "expert_trajectories") and env.expert_trajectories:
|
||||
return env.expert_trajectories
|
||||
if not hasattr(env, "engine") or not hasattr(env.engine, "traffic_manager"):
|
||||
return {}
|
||||
from metadrive.type import MetaDriveType
|
||||
data = getattr(env.engine.traffic_manager, "current_traffic_data", None)
|
||||
if not data:
|
||||
return {}
|
||||
expert_trajs = {}
|
||||
for scenario_id, track in data.items():
|
||||
if track.get("type") != MetaDriveType.VEHICLE or "state" not in track:
|
||||
continue
|
||||
state = track["state"]
|
||||
positions = state.get("position")
|
||||
if positions is None:
|
||||
continue
|
||||
valid = state.get("valid", np.ones(len(positions), dtype=bool))
|
||||
valid = np.asarray(valid).flatten()
|
||||
if valid.size != len(positions):
|
||||
valid = np.ones(len(positions), dtype=bool)
|
||||
first_show = int(np.argmax(valid)) if valid.any() else 0
|
||||
last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else len(positions) - 1
|
||||
obj_id = track.get("metadata", {}).get("object_id", str(scenario_id))
|
||||
expert_trajs[obj_id] = {
|
||||
"positions": np.asarray(positions),
|
||||
"start_timestep": first_show,
|
||||
"end_timestep": last_show,
|
||||
}
|
||||
return expert_trajs
|
||||
|
||||
|
||||
def _run_trajectory_animation(expert_trajs, scenario_idx):
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.animation import FuncAnimation
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
print("No expert trajectories to visualize.")
|
||||
return
|
||||
|
||||
fig, ax = plt.subplots(figsize=(12, 12))
|
||||
max_timestep = max(t["end_timestep"] for t in expert_trajs.values())
|
||||
min_timestep = min(t["start_timestep"] for t in expert_trajs.values())
|
||||
|
||||
colors = plt.cm.tab10(np.linspace(0, 1, len(expert_trajs)))
|
||||
for idx, (obj_id, traj) in enumerate(expert_trajs.items()):
|
||||
positions = np.asarray(traj["positions"])
|
||||
if positions.ndim >= 2:
|
||||
positions = positions[:, :2]
|
||||
else:
|
||||
continue
|
||||
ax.plot(
|
||||
positions[:, 0], positions[:, 1],
|
||||
color=colors[idx], alpha=0.3, linewidth=1,
|
||||
label=f"Vehicle {str(obj_id)[:6]}",
|
||||
)
|
||||
|
||||
scatter = ax.scatter([], [], s=200, c="red", marker="o", edgecolors="black", linewidths=2)
|
||||
time_text = ax.text(0.02, 0.95, "", transform=ax.transAxes, fontsize=14)
|
||||
ax.set_xlabel("X (m)")
|
||||
ax.set_ylabel("Y (m)")
|
||||
ax.set_title(f"Expert Trajectory Visualization - Scenario {scenario_idx}")
|
||||
ax.legend(loc="upper right", fontsize=8)
|
||||
ax.grid(True, alpha=0.3)
|
||||
ax.axis("equal")
|
||||
|
||||
def update(frame):
|
||||
current_time = min_timestep + frame
|
||||
current_positions = []
|
||||
for traj in expert_trajs.values():
|
||||
st, et = traj["start_timestep"], traj["end_timestep"]
|
||||
if st <= current_time <= et:
|
||||
pos = np.asarray(traj["positions"])
|
||||
if pos.ndim >= 2:
|
||||
pos = pos[current_time - st, :2]
|
||||
else:
|
||||
continue
|
||||
current_positions.append(pos)
|
||||
if current_positions:
|
||||
scatter.set_offsets(np.array(current_positions))
|
||||
time_text.set_text(f"Time: {frame * 0.1:.1f}s (Frame {frame})")
|
||||
return scatter, time_text
|
||||
|
||||
anim = FuncAnimation(
|
||||
fig, update, frames=max_timestep - min_timestep + 1,
|
||||
interval=100, blit=True, repeat=True,
|
||||
)
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
return anim
|
||||
|
||||
|
||||
def _run_trajectory(args):
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
data_dir = _resolve_data_dir(args.data_dir)
|
||||
data_path = os.path.abspath(data_dir)
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": 500,
|
||||
"use_render": False,
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.scenario_idx,
|
||||
"num_scenarios": 1,
|
||||
"log_level": 40,
|
||||
}
|
||||
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
try:
|
||||
env.reset(seed=args.scenario_idx)
|
||||
expert_trajs = _build_expert_trajectories_from_env(env)
|
||||
_run_trajectory_animation(expert_trajs, args.scenario_idx)
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
# --- Main ---
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Unified visualization: replay, policy (BC/MAGAIL), trajectory.",
|
||||
)
|
||||
subparsers = parser.add_subparsers(dest="mode", required=True, help="replay | policy | trajectory")
|
||||
|
||||
# Common args for data_dir (used by all)
|
||||
def add_common_data_args(p):
|
||||
p.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Waymo scenario directory")
|
||||
p.add_argument("--start_index", type=int, default=0)
|
||||
p.add_argument("--num_scenarios", type=int, default=1)
|
||||
p.add_argument("--horizon", type=int, default=200)
|
||||
|
||||
# replay
|
||||
pr = subparsers.add_parser("replay", help="Replay scenario with ExpertReplayEnv (no policy)")
|
||||
add_common_data_args(pr)
|
||||
pr.set_defaults(horizon=500)
|
||||
|
||||
# policy
|
||||
pp = subparsers.add_parser("policy", help="Visualize BC or MAGAIL trained policy")
|
||||
add_common_data_args(pp)
|
||||
pp.add_argument("--policy_type", type=str, default="auto", choices=["auto", "bc", "magail"])
|
||||
pp.add_argument("--model_path", type=str, default="models/bc/policy_best.pt")
|
||||
pp.add_argument("--deterministic", action="store_true", help="MAGAIL: use mean action")
|
||||
|
||||
# trajectory
|
||||
pt = subparsers.add_parser("trajectory", help="2D matplotlib animation of expert trajectories")
|
||||
pt.add_argument("--data_dir", type=str, default="data/exp_filtered")
|
||||
pt.add_argument("--scenario_idx", type=int, default=0)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Resolve data_dir relative to project root when default
|
||||
if args.mode != "trajectory":
|
||||
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||
else:
|
||||
if args.data_dir in ("data/exp_filtered", "data/exp_converted"):
|
||||
args.data_dir = os.path.join(project_root, args.data_dir)
|
||||
|
||||
if args.mode == "replay":
|
||||
_run_replay(args)
|
||||
elif args.mode == "policy":
|
||||
_run_policy(args)
|
||||
elif args.mode == "trajectory":
|
||||
_run_trajectory(args)
|
||||
else:
|
||||
parser.error(f"Unknown mode: {args.mode}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,105 +0,0 @@
|
||||
import sys
|
||||
import os
|
||||
|
||||
# 添加路径
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
project_root = os.path.dirname(current_dir)
|
||||
env_dir = os.path.join(project_root, "Env")
|
||||
sys.path.insert(0, project_root)
|
||||
sys.path.insert(0, env_dir)
|
||||
|
||||
# 现在可以导入了
|
||||
from scenario_env import MultiAgentScenarioEnv
|
||||
from metadrive.engine.asset_loader import AssetLoader
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.animation import FuncAnimation
|
||||
|
||||
class DummyPolicy:
|
||||
"""
|
||||
占位策略,用于数据检查时初始化环境
|
||||
不需要实际执行动作,只是为了满足环境初始化要求
|
||||
"""
|
||||
def act(self, *args, **kwargs):
|
||||
# 返回零动作 [throttle, steering]
|
||||
return np.array([0.0, 0.0])
|
||||
|
||||
def visualize_expert_trajectory(env, scenario_idx=0):
|
||||
"""
|
||||
可视化专家轨迹的俯视图动画
|
||||
"""
|
||||
env.reset()
|
||||
expert_trajs = env.expert_trajectories
|
||||
|
||||
if len(expert_trajs) == 0:
|
||||
print("当前场景无专家轨迹")
|
||||
return
|
||||
|
||||
# 设置绘图
|
||||
fig, ax = plt.subplots(figsize=(12, 12))
|
||||
|
||||
# 获取所有轨迹的最大时间长度
|
||||
max_timestep = max(traj["end_timestep"] for traj in expert_trajs.values())
|
||||
min_timestep = min(traj["start_timestep"] for traj in expert_trajs.values())
|
||||
|
||||
# 绘制完整轨迹(淡色)
|
||||
colors = plt.cm.tab10(np.linspace(0, 1, len(expert_trajs)))
|
||||
for idx, (obj_id, traj) in enumerate(expert_trajs.items()):
|
||||
positions = traj["positions"][:, :2]
|
||||
ax.plot(positions[:, 0], positions[:, 1],
|
||||
color=colors[idx], alpha=0.3, linewidth=1,
|
||||
label=f'Vehicle {obj_id[:6]}')
|
||||
|
||||
# 初始化当前位置标记
|
||||
scatter = ax.scatter([], [], s=200, c='red', marker='o', edgecolors='black', linewidths=2)
|
||||
time_text = ax.text(0.02, 0.95, '', transform=ax.transAxes, fontsize=14)
|
||||
|
||||
ax.set_xlabel('X (m)')
|
||||
ax.set_ylabel('Y (m)')
|
||||
ax.set_title(f'Expert Trajectory Visualization - Scenario {scenario_idx}')
|
||||
ax.legend(loc='upper right', fontsize=8)
|
||||
ax.grid(True, alpha=0.3)
|
||||
ax.axis('equal')
|
||||
|
||||
def update(frame):
|
||||
current_time = min_timestep + frame
|
||||
|
||||
# 收集当前时间所有车辆的位置
|
||||
current_positions = []
|
||||
for traj in expert_trajs.values():
|
||||
if traj["start_timestep"] <= current_time <= traj["end_timestep"]:
|
||||
idx = current_time - traj["start_timestep"]
|
||||
pos = traj["positions"][idx, :2]
|
||||
current_positions.append(pos)
|
||||
|
||||
if len(current_positions) > 0:
|
||||
current_positions = np.array(current_positions)
|
||||
scatter.set_offsets(current_positions)
|
||||
|
||||
time_text.set_text(f'Time: {frame * 0.1:.1f}s (Frame {frame})')
|
||||
return scatter, time_text
|
||||
|
||||
anim = FuncAnimation(fig, update, frames=max_timestep-min_timestep+1,
|
||||
interval=100, blit=True, repeat=True)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
|
||||
return anim
|
||||
|
||||
if __name__ == "__main__":
|
||||
WAYMO_DATA_DIR = r"/home/huangfukk/mdsn"
|
||||
data_dir = AssetLoader.file_path(WAYMO_DATA_DIR, "exp_filtered", unix_style=False)
|
||||
|
||||
env = MultiAgentScenarioEnv(
|
||||
config={
|
||||
"data_directory": data_dir,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 3,
|
||||
"use_render": False,
|
||||
},
|
||||
agent2policy=DummyPolicy()
|
||||
)
|
||||
|
||||
# 可视化第一个场景
|
||||
anim = visualize_expert_trajectory(env, scenario_idx=0)
|
||||
@@ -1,93 +0,0 @@
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
# Add project root to Python path so we can import Env module
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
from Env.expert_replay_env import ExpertReplayEnv
|
||||
|
||||
def visualize_replay(args):
|
||||
data_path = os.path.abspath(args.data_dir)
|
||||
if not os.path.exists(data_path):
|
||||
raise ValueError(f"Data directory {data_path} not found")
|
||||
|
||||
# Same as data generation: avoid MetaDrive assertion when requested num_scenarios > available.
|
||||
from metadrive.scenario.utils import read_dataset_summary
|
||||
_, summary_lookup, _ = read_dataset_summary(data_path)
|
||||
if args.start_index >= len(summary_lookup):
|
||||
raise ValueError(
|
||||
f"start_index={args.start_index} out of range. Dataset has {len(summary_lookup)} scenarios."
|
||||
)
|
||||
max_available = len(summary_lookup) - args.start_index
|
||||
num_to_run = min(args.num_scenarios, max_available)
|
||||
|
||||
env_config = {
|
||||
"data_directory": data_path,
|
||||
"is_multi_agent": True,
|
||||
"num_controlled_agents": 100,
|
||||
"horizon": args.horizon,
|
||||
"use_render": True, # Enable rendering
|
||||
"sequential_seed": True,
|
||||
"reactive_traffic": False,
|
||||
"start_scenario_index": args.start_index,
|
||||
"num_scenarios": -1,
|
||||
"log_level": 40, # ERROR
|
||||
# "pstats": True, # For performance debugging
|
||||
}
|
||||
|
||||
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||
env = ExpertReplayEnv(config=env_config)
|
||||
|
||||
try:
|
||||
for i in range(args.start_index, args.start_index + num_to_run):
|
||||
print(f"\n--- Playing Scenario {i} ---")
|
||||
try:
|
||||
obs = env.reset(seed=i)
|
||||
except Exception as e:
|
||||
print(f"Error resetting scenario {i}: {e}")
|
||||
continue
|
||||
|
||||
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
||||
|
||||
for step in range(args.horizon):
|
||||
# Step
|
||||
obs, rewards, dones, infos = env.step(None)
|
||||
|
||||
# Render
|
||||
env.render(mode="top_down",
|
||||
text={
|
||||
"Step": step,
|
||||
"Agents": len(env.controlled_agents),
|
||||
"Scenario": i
|
||||
})
|
||||
|
||||
# Sleep to control playback speed
|
||||
time.sleep(0.05)
|
||||
|
||||
if dones["__all__"]:
|
||||
print(f"Scenario {i} finished at step {step}")
|
||||
break
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("Interrupted by user")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
print(f"Global error: {e}")
|
||||
finally:
|
||||
env.close()
|
||||
print("Environment closed.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/exp_filtered", help="Path to Waymo data")
|
||||
parser.add_argument("--start_index", type=int, default=0)
|
||||
parser.add_argument("--num_scenarios", type=int, default=1)
|
||||
parser.add_argument("--horizon", type=int, default=500)
|
||||
|
||||
args = parser.parse_args()
|
||||
visualize_replay(args)
|
||||
@@ -1,189 +0,0 @@
|
||||
"""
|
||||
Unified visualization for BC and MAGAIL trained policies.
|
||||
Use --policy_type bc or magail (or auto-detect from --model_path: .pt -> bc, else magail).
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
# Add project root to Python path
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
from Env.bc_env import BCScenarioEnv
|
||||
from metadrive.engine.engine_utils import close_engine
|
||||
|
||||
|
||||
def _resolve_data_dir(args):
|
||||
"""Resolve data directory: explicit or auto-detect under project data/."""
|
||||
if args.data_dir:
|
||||
data_dir = args.data_dir
|
||||
else:
|
||||
current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
data_dir = os.path.join(current_dir, "data", "exp_filtered")
|
||||
if not os.path.exists(data_dir):
|
||||
data_dir = os.path.join(current_dir, "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):
|
||||
"""Resolve model path: if not found, try models/bc or models/magail."""
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
if policy_type == "bc":
|
||||
candidate = os.path.join("models", "bc", model_path)
|
||||
else:
|
||||
candidate = os.path.join("models", "magail", model_path)
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
if policy_type == "magail" and not model_path.endswith("_actor.pth"):
|
||||
candidate = model_path + "_actor.pth"
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
raise FileNotFoundError(f"Model path {model_path} not found (tried {candidate}).")
|
||||
|
||||
|
||||
def visualize_model(args):
|
||||
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_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:
|
||||
actions = {}
|
||||
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()
|
||||
|
||||
for idx, aid in enumerate(agent_ids):
|
||||
actions[aid] = actions_np[idx].flatten()
|
||||
|
||||
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()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Visualize BC or MAGAIL trained policy in 45-dim scenario env."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--policy_type",
|
||||
type=str,
|
||||
default="auto",
|
||||
choices=["auto", "bc", "magail"],
|
||||
help="Policy type: bc (StateIndependentPolicy .pt) or magail (Actor _actor.pth). auto = infer from model_path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model_path",
|
||||
type=str,
|
||||
default="models/bc/policy_best.pt",
|
||||
help="Path to model: BC .pt (e.g. models/bc/policy_best.pt) or MAGAIL _actor.pth (e.g. models/magail/model_50_actor.pth)",
|
||||
)
|
||||
parser.add_argument("--data_dir", type=str, default=None, help="Waymo data directory (default: data/exp_filtered)")
|
||||
parser.add_argument("--start_index", type=int, default=0)
|
||||
parser.add_argument("--num_scenarios", type=int, default=1)
|
||||
parser.add_argument("--horizon", type=int, default=200)
|
||||
parser.add_argument("--deterministic", action="store_true", help="For MAGAIL: use mean action instead of sampling")
|
||||
|
||||
args = parser.parse_args()
|
||||
visualize_model(args)
|
||||
Reference in New Issue
Block a user