From 21c046aef0ace492c11e9bd5fd67f9a703ecd42f Mon Sep 17 00:00:00 2001 From: huangfu <3045324663@qq.com> Date: Mon, 2 Feb 2026 01:18:18 +0800 Subject: [PATCH] =?UTF-8?q?BC=E7=AE=97=E6=B3=95=E5=AE=9E=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Algorithm/bc.py | 52 +++++ Env/__pycache__/scenario_env.cpython-313.pyc | Bin 12272 -> 12915 bytes Env/__pycache__/scenario_env.cpython-39.pyc | Bin 6801 -> 6958 bytes Env/bc_env.py | 64 ++++++ Env/scenario_env.py | 27 +++ README.md | 150 +++++++------- dataset/expert_dataset.py | 3 +- ...vents.out.tfevents.1769858761.Hfkk.41584.0 | Bin 0 -> 5884 bytes ...vents.out.tfevents.1769859823.Hfkk.44173.0 | Bin 0 -> 1532 bytes ...vents.out.tfevents.1769860931.Hfkk.47961.0 | Bin 0 -> 227 bytes ...vents.out.tfevents.1769862091.Hfkk.51799.0 | Bin 0 -> 88 bytes ...vents.out.tfevents.1769862096.Hfkk.51875.0 | Bin 0 -> 1581 bytes ...vents.out.tfevents.1769862459.Hfkk.53251.0 | Bin 0 -> 1532 bytes ...vents.out.tfevents.1769862828.Hfkk.54652.0 | Bin 0 -> 1532 bytes ...vents.out.tfevents.1769863419.Hfkk.56732.0 | Bin 0 -> 3370 bytes ...vents.out.tfevents.1769864090.Hfkk.59264.0 | Bin 0 -> 88 bytes ...vents.out.tfevents.1769875401.Hfkk.64137.0 | Bin 0 -> 15072 bytes ...ents.out.tfevents.1769965130.Hfkk.240848.0 | Bin 0 -> 3080 bytes scripts/README.md | 82 ++++++++ scripts/README_visualize.md | 113 ----------- scripts/generate_expert_data.py | 4 +- scripts/visualize_trained_policy.py | 186 +++++++++++------- train_bc.py | 186 ++++++++++++++++++ train_magail.py | 70 +------ visualize_bc.py | 17 ++ 25 files changed, 632 insertions(+), 322 deletions(-) create mode 100644 Algorithm/bc.py create mode 100644 Env/bc_env.py create mode 100644 logs/20260131-192601/events.out.tfevents.1769858761.Hfkk.41584.0 create mode 100644 logs/20260131-194343/events.out.tfevents.1769859823.Hfkk.44173.0 create mode 100644 logs/20260131-200211/events.out.tfevents.1769860931.Hfkk.47961.0 create mode 100644 logs/20260131-202131/events.out.tfevents.1769862091.Hfkk.51799.0 create mode 100644 logs/20260131-202136/events.out.tfevents.1769862096.Hfkk.51875.0 create mode 100644 logs/20260131-202739/events.out.tfevents.1769862459.Hfkk.53251.0 create mode 100644 logs/20260131-203348/events.out.tfevents.1769862828.Hfkk.54652.0 create mode 100644 logs/20260131-204339/events.out.tfevents.1769863419.Hfkk.56732.0 create mode 100644 logs/20260131-205450/events.out.tfevents.1769864090.Hfkk.59264.0 create mode 100644 logs/20260201-000321/events.out.tfevents.1769875401.Hfkk.64137.0 create mode 100644 logs/bc/20260202-005850/events.out.tfevents.1769965130.Hfkk.240848.0 create mode 100644 scripts/README.md delete mode 100644 scripts/README_visualize.md create mode 100644 train_bc.py create mode 100644 visualize_bc.py diff --git a/Algorithm/bc.py b/Algorithm/bc.py new file mode 100644 index 0000000..b2695b2 --- /dev/null +++ b/Algorithm/bc.py @@ -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 diff --git a/Env/__pycache__/scenario_env.cpython-313.pyc b/Env/__pycache__/scenario_env.cpython-313.pyc index 363686e31ffca5585eebd371b4e52d5a4fdfadbc..5bb51a4e735e6466b9dc969194142791421fc777 100644 GIT binary patch delta 1259 zcma)6Uu;WJ7(b`hPSugO;t!g%g$g&^=K}>+?|i@S^rW8>UsBay z9F8g?KI5UY>35;~)!*%=rPiX|@vOoWZ-+v)zqVUyH=%-%Witz5Qv{B?7B;Mz-OpKm zO70M%*Vixmi>kZkk0T@ck#r5z1l-W&S~b)C@71n zMDB^pZQ8r#w-9xE=r(Lkx9hkQrp&dWkTF<{?A>GVUu#d85X|V7V=f=Gx)&WYE`UTu zF4JtNQX`jw=o+|2wQ=YNanuGPpSv4`hvVw5UsSL8=tei^Zdiv$0s%XvxORDK1*$6bZ zfmHj^UwM=%OjSAW7p?g*0h&3Pks0S1!Qp%DASesoTP|Rt4g$6XHZCi&x=HMntHrb4 zI@7u70I|DA+r~HKHx!&@TeM`0PVRdYzaKBR?-JkXY9e_@)$!B$(`9e89l04(Kysa!Xm!#&hw5}wrD@&avsq>}OHEsFdD(Bp?H}mbq8;ip88i#~IegHv5 z>vsJ*W@qP4j9-8cOB&4j+_1#A(Q+Mi5f}}6VaVsD6UD2(7nap0F>39T)r6ulRpys* zZy92l6SAUJWLQ?#9D}*ExcOW*fm^P^SW|-DgXyM@rh-uOTL?!4XbzRL-rjUFkxBO* zRry0q+=V^C^%UT8FhnPd#UL|>u_wnD$2$;k06ro>t75&vf^Z@aaY|Pr$`G3 zYZ0CZ(B`ppN~`4ZFykqOUxD*sKOKRa;d=VA_&6M+P5K4lW1i`B1z8^98Ul{V2cV;6 zd+R7t{QYNJBC{gk9}k~G5D%4TmD!YdHS8oGfd?%s=m5NJsrPqNbK^v?P*F%0tfj@v uzYz?4((!Y@3F;zS=^02yZqpphUwM(fhWyH1HNz&ld}7B}f-Colvwr}J*)-Dt delta 833 zcmYjPK}-`-5Zzz8C8aDPTU($?TUydmw3Z^FF+of`7&VeeEFmJU-*&;)THtrB(L_uy zC_y1Qns89$;>nZQgK|-0JbBW1F?;f+7ZU(Twl4AyZ{N(mnK$$EH?nI=!-gc) z3H;oe&E{5v&l`T%+w!$~z4YD=>rGzL3r~Wr@J1F%pUJ++E%2hL8x~|QjLB|t2=2AZ z zSBaNaShD0aSqNZG?MNSwz^7lIyuwfh;qh9#Qg*Wtuz$Zsf8MDp_ zRZ~hD)l@c%v7<~2SJaYWFVLcDNK=zJ-fOa`V90&k5rsShEs_d9M&o6K^H?WojE;}z zGD;ygaZzJB{&)%>A_vJLxS|pAQ16LSu?sy}8X<)+jZooW)X{N10-r>kVmK~agxCAp zLhq28R(;6cqO^wQvPLaauTqvRvAO&4Hc|n9+)q~Zfq0sP&3hR}xoI?q>@LDQ0&dCj z@N)n0h)d5T4oWX8(8dN8J2K5)ud-kN^QhN-Q9RBqbsSu*wHGdF;I1on>}r?Dvwe zRIEL4@QO@XyD6L2&(mM$&O~+QFGoia%h7KT$6|$>ENgSAjCEnVB}+9}OnO3?EERe~ zJBU{WH!Y{adO~M=WSADP{?Joe0}zffcu85V6LfRW6OPHYAH@dr?n#UQrnbN6{!Z6t z{d(dzea^p;{7~QA{vr8^-c-YzS}kq%JcvBh;ThSr5wcqT+KAa3GEZ7eZFKA9$C|F~ zvVG65tL<&EG@Mt@*V-K{gu}%H49v*(Y9mL#Y?V48Xf0Vi1 z$CB=ynkXs5GGQ;}X54o(qgTh#47xCFL6%o2nRV`<0YP4Re8Ul@(OM>he3<8T(Qy<{ zz^k;l+jc9s0I23L=YO7=%&?@*?>7v~u@v)v%DmoWc?G;&wo0At*+8dw9m3>_2Qf>D zx@382mJC@E4hPGnfHSx^i@&?Dv^;-#0Vcwjxzh+Dj)!H3N}b;ex@_SDU7qWqaJb}3 z0bHVx7B#J@1+kJXIKUDuje(Ftv26G*C{H06F^Ty@n#ii@U;>yCp6*x!Cqy9aIS?`PV%ECoB zgd`T-GCAx};5)QgG0Ob&{d<}$bB9ZTAGh%GLDs7_q*amz3dS~BR$5qjQz-047j~#` zb`l7zk2GZ17D9)GF(iU)=i8^MLh<4toI-Ew;_ncjtO%^G;!40R9A6H)*SWnWWDunt zS=Mt=UWV%!G;C!zYg?7-Znv{uY#F1|iXnx)hyuRH{E2~)CQF4Ugq^|?o`90T^g;~} zV_CD|@H)13>vkE6qcg$)EUBQ6!lmjIQ(+~bL=KbSECS2Ea58vAl!-`%sj~fV2WDa% z9baw#J|K1dod5oyrT^kT8T?>!ohFAO{#VHj^gSj2eS#zYt(O*?|2OYhXI?Sg5+7x> zmD}(ygtQ0@fu(pE6O4Hl9;}j=B4U|ND~J9$bQKe=ORmwJ@TU7T?WhI{$3eW^@o<|) znB%f5{~JD?_M^X2Sh};Hwm3jAP1DmQY+hLxs_L3BOYR#Yg545_&|>B{PZ98;L?<3WzHff|#>b_Zs0A%%2~c w$G`E9h9(d{8#<%E=|3GhjCSwvzEPIL__ztEk0YvqQlW2mt9@tuw}&tM1D~u>G5`Po delta 1851 zcmY*ZU2GIp6rL$N+x?qumzJ_!Dy2|6=ue=vE(Fa0e`XVoeM~zR0_;MeNJ{Ut{;)5?5Up!}~tfsKzJeJl~DSgPV!buc#dBA#OgJZqgVm#Te+8f zXqrbh~E3!!Z34(aH0Q1=&DD4RlCS42&E zkyCHL*2@RsI0ovZ>(^=mfUW_?)B$@(PN!XdZIMY&hFZOCk8RXlA@FKdR#=mJ9HX_m zXAcbNuE43VpdvgD(x>xtl(F-|#MJco%mkFxV|%b|fe1vX?dZV*;RQYfx-;a#)zdZ5 zi69g$oo({gdh+brWuZ~$z~>J0M7ZF?Rc+OPx9Zw*xhDc$2>ljw`~{EU-8urWFstGU zjt|Gj07x3fVMAM0*XN<`5-l0{j=QY8=rc<1v9K-!Rt1i8%|6k;ws20Bt1wBQWU3NSwR8hC&{Yb3l(R_CPv9WX3Bd&f z!%UhPGi45&sg9JHQ6F}VsE1vDw@;%>XDhxJg73P*VK9RhI71NK_6D7w7j=onF^sCp z+7n9X_@BdB{GyzPOZ4G!f{O?`7qmFeFqRS6WT&#AL{Kn{VFR`hjVD8_jG%2)0&di* zB-FV=AMR_tu-gjJdwu28!f$4K8`8%e6y0-v_Qa`E)8lMrdUEPKo1C4PI0Mt<>3!OB zX<}yP>_uI8+B!S&GCW7omMa|&k6)q0TM^0$o$=wfBR(iJ2d=-KjYedbvxas-rs!A{6d{5_H=!T z#_#b5UlF{Z{@FYkd)nmrVa-#{#ww!`#={&XMxuing3jb;qYVW700|W;VaD`AMK~scv-+?1! z?k4CY0D?&N>xDpyY-xgJf_D&fr-$+bz8EzC&yqdb<_MzYj`sQxIu4E?1bKv%* zxzxbXR53LSWW{cAB&5Mrjny_eE KL+X#x$^QUE3@wEK diff --git a/Env/bc_env.py b/Env/bc_env.py new file mode 100644 index 0000000..b802fba --- /dev/null +++ b/Env/bc_env.py @@ -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 diff --git a/Env/scenario_env.py b/Env/scenario_env.py index 272b6bb..e7cad7f 100644 --- a/Env/scenario_env.py +++ b/Env/scenario_env.py @@ -100,6 +100,33 @@ class MultiAgentScenarioEnv(ScenarioEnv): for scenario_id in _obj_to_clean_this_frame: self.engine.traffic_manager.current_traffic_data.pop(scenario_id) + # Fix: Ensure all objects are cleared properly before reset + # Instead of manually clearing, we just let the engine handle it, but we might need to + # ensure no stale references in managers. + + # The error "KeyError" in clear_objects usually means we are trying to clear an object + # that is already gone from _spawned_objects but still tracked by a manager. + + # Try to clear only objects that actually exist in the engine + # existing_objects = list(self.engine.get_objects().keys()) + # if existing_objects: + # self.engine.clear_objects(existing_objects) + + # Force clear agent manager's spawned objects to avoid stale references + if hasattr(self.engine, 'agent_manager') and self.engine.agent_manager: + # Check if it's ScenarioAgentManager or VehicleAgentManager + # ScenarioAgentManager might not have spawned_objects directly exposed or named differently + # But BaseAgentManager usually has it. + # If it's ScenarioAgentManager, it might be using a different structure. + + # Safe clear for BaseAgentManager subclasses + if hasattr(self.engine.agent_manager, 'spawned_objects'): + self.engine.agent_manager.spawned_objects.clear() + + # Also clear active_objects if present (VehicleAgentManager uses this) + if hasattr(self.engine.agent_manager, '_active_objects'): + self.engine.agent_manager._active_objects.clear() + self.engine.reset() self.reset_sensors() self.engine.taskMgr.step() diff --git a/README.md b/README.md index 854f995..5e0f964 100644 --- a/README.md +++ b/README.md @@ -1,98 +1,96 @@ # MAGAIL4AutoDrive -> 基于多智能体生成对抗模仿学习(MAGAIL)的自动驾驶训练系统 | MetaDrive + Waymo Open Motion Dataset +基于 **MetaDrive** 仿真器和 **Waymo Open Motion Dataset** 的自动驾驶多智能体模仿学习(MAGAIL)与行为克隆(BC)训练系统。 -本项目利用 Waymo 真实驾驶数据,通过 MetaDrive 仿真环境构建专家回放系统,提取车辆状态与动作,用于训练多智能体模仿学习算法 (MAGAIL)。 +本项目旨在从真实的 Waymo 驾驶数据中提取专家轨迹,并通过模仿学习(Imitation Learning)训练能够适应复杂交互场景的自动驾驶策略。 -## 📁 核心模块 +## 目录结构 -* **`Env/expert_replay_env.py`**: 专家回放环境。核心类 `ExpertReplayEnv`,负责读取 Waymo 轨迹,计算逆动力学动作,并过滤非道路/静态车辆。 -* **`Env/inverse_dynamics.py`**: 逆动力学模块。根据车辆位置和航向计算油门、刹车和转向动作。 -* **`scripts/generate_expert_data.py`**: 数据收集脚本。批量运行场景并保存训练数据。 -* **`scripts/visualize_replay.py`**: 可视化脚本。用于观察回放效果和数据质量。 - -*** - -## 🚀 1. 数据收集 - -### 生成专家数据 -使用 `generate_expert_data.py` 脚本从 Waymo 数据集中批量提取 (State, Action) 对。 - -```bash -# 设置 Python 路径 -export PYTHONPATH=$PYTHONPATH:.:./metadrive - -# 运行生成脚本 -# --data_dir: Waymo 数据路径 (建议使用 exp_filtered) -# --output_dir: 结果保存路径 -# --num_scenarios: 要处理的场景数量 -python scripts/generate_expert_data.py \ - --data_dir data/exp_filtered \ - --output_dir data/training_data \ - --num_scenarios 100 \ - --start_index 0 +```text +MAGAIL4AutoDrive/ +├── Algorithm/ # 强化学习与模仿学习算法实现 +│ ├── policy.py # 基础策略网络 (MLP 等) +│ ├── ppo.py # PPO 算法实现 +│ ├── magail.py # MAGAIL 算法核心逻辑 +│ ├── disc.py # 判别器 (Discriminator) 网络 +│ └── ... +├── Env/ # 仿真环境封装 (MetaDrive Wrapper) +│ ├── bc_env.py # BCScenarioEnv,45 维观测(BC/MAGAIL 共用) +│ ├── scenario_env.py # 多智能体基础场景环境 +│ ├── expert_replay_env.py # 专家轨迹回放环境(数据生成与回放) +│ ├── inverse_dynamics.py # 逆动力学模块 (轨迹 -> 动作) +│ ├── simple_idm_policy.py # ConstantVelocityPolicy 占位策略 +│ └── ... +├── dataset/ # 数据集加载器 +│ ├── expert_dataset.py # 通用专家数据加载类 +│ └── magail_dataset.py # MAGAIL 训练专用数据加载器 +├── scripts/ # 工具脚本(数据、回放、可视化、分析) +│ ├── generate_expert_data.py # 从 Waymo 生成专家 (obs, act) pkl +│ ├── visualize_replay.py # 原始专家数据回放 +│ ├── visualize_trained_policy.py # BC/MAGAIL 策略可视化统一入口 +│ ├── 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 训练 +├── visualize_bc.py # [根目录] BC 可视化薄包装 -> scripts/visualize_trained_policy.py +└── README.md ``` -**生成的 `.pkl` 文件结构**: -包含一个列表,每个元素是一条车辆轨迹(Trajectory Dictionary): -* `obs`: `(T, 45)` - 观测矩阵。包含 Ego 状态 (5维) + 10辆邻居车相对信息 (40维)。 -* `acts`: `(T, 2)` - 动作矩阵。`[Steering, Accel]`,归一化到 `[-1, 1]`。 -* `agent_id`: 车辆 ID。 -* `scenario_id`: 所属场景 ID。 +## 路径约定(相对项目根) -**内置过滤器**: -脚本会自动过滤掉以下无效车辆: -1. **非道路车辆**:始终在停车场或路外行驶的车辆。 -2. **静态车辆**:全称移动距离小于 5米 且速度从未超过 1m/s 的车辆(作为背景流存在,不收集数据)。 +- **数据**:Waymo 场景 `data/exp_filtered`;专家 pkl `data/training_data`;其他轨迹 `data/trajectories` +- **模型**:BC `models/bc/`,MAGAIL `models/magail/` +- **日志**:TensorBoard 写入 `logs/bc/`、`logs/magail/` ---- +所有默认路径均为相对项目根,便于在不同设备上复用。 -## 🔍 2. 数据可视化与验证 +## 核心工作流 -### 回放可视化 -使用 `visualize_replay.py` 直观地观察回放效果,确认车辆行为是否自然,以及过滤逻辑是否生效。 +### 1. 数据准备 +使用 `scripts/generate_expert_data.py` 将 Waymo 数据转换为训练用 `.pkl`,输出到 `data/training_data/`。 ```bash -# 运行可视化 -# --horizon: 回放的最大步数 (Waymo 场景通常为 90 或 198 步) -python scripts/visualize_replay.py \ - --data_dir data/exp_filtered \ - --start_index 0 \ - --num_scenarios 1 \ - --horizon 200 +python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 ``` -**观察要点**: -* **受控车辆 (Controlled Agents)**:控制台会显示数量(如 `Controlled agents: 2`)。这些是真正产生数据的车辆。 -* **背景车辆**:如果在渲染图中看到其他车(通常是路边停放的),但受控数量很少,说明静态过滤生效了。 +### 2. 行为克隆 (BC) +- **训练**:`python train_bc.py`(模型保存到 `models/bc/`,日志到 `logs/bc/`) +- **可视化**:`python visualize_bc.py` 或 `python scripts/visualize_trained_policy.py --policy_type bc --model_path models/bc/policy_best.pt` -### 数据分析 -使用 `analyze_expert_data.py` 查看生成数据的统计分布。 +### 3. 多智能体对抗模仿学习 (MAGAIL) +- **训练**:`python train_magail.py`(模型保存到 `models/magail/`,日志到 `logs/magail/`) +- **可视化**:`python scripts/visualize_trained_policy.py --policy_type magail --model_path models/magail/model_50_actor.pth` -```bash -python scripts/analyze_expert_data.py --data_path data/training_data/expert_data_0_100.pkl -``` +### 4. 策略可视化统一入口 +BC 与 MAGAIL 共用 `scripts/visualize_trained_policy.py`,通过 `--policy_type bc|magail`(或根据 `--model_path` 自动推断)选择模型类型。根目录 `visualize_bc.py` 为 BC 的薄包装。详见 [scripts/README_visualize.md](scripts/README_visualize.md) 与 [scripts/README.md](scripts/README.md)。 ---- +## 文件与模块职责 -## 🧠 3. 模型训练 (Next Steps) +### 根目录脚本 +- **train_bc.py**:BC 训练,加载 `data/training_data` 下 pkl,模型与日志写入 `models/bc/`、`logs/bc/` +- **train_magail.py**:MAGAIL 训练,环境使用 `BCScenarioEnv`(45 维),模型与日志写入 `models/magail/`、`logs/magail/` +- **visualize_bc.py**:薄包装,调用 `scripts/visualize_trained_policy.py --policy_type bc` -有了 `data/training_data/` 下的专家数据后,您可以开始训练 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**:轨迹 → 油门/转向动作 -### 训练流程 -1. **加载数据**:使用 `dataset/expert_dataset.py` 中的 `ExpertDataset` 类加载 `.pkl` 数据。 -2. **初始化 MAGAIL**: - * **Generator (Policy)**: 接收观测 `(B, 45)`,输出动作 `(B, 2)`。 - * **Discriminator**: 接收状态-动作对 `(s, a)`,判断是专家还是生成器。 -3. **交互采样**: - * 在 `MultiAgentScenarioEnv`(非回放模式)中运行 Policy。 - * 收集 Policy 生成的轨迹。 -4. **对抗更新**: - * 利用专家数据和 Policy 数据训练 Discriminator。 - * 利用 Discriminator 的输出作为 Reward (GAIL Reward) 训练 Policy (PPO/TRPO)。 +### Algorithm 模块 +- **Algorithm/policy.py**:`StateIndependentPolicy`,BC 使用的 MLP 策略 -### 推荐配置 -* **Observation**: 45维 (Ego + 10 Neighbors) -* **Action**: 2维 Continuous (Steering, Accel) -* **Horizon**: 200 steps -* **Batch Size**: 1024+ (多智能体环境下数据量很大) +### scripts 目录 +工具脚本用途与用法见 [scripts/README.md](scripts/README.md)。 diff --git a/dataset/expert_dataset.py b/dataset/expert_dataset.py index 1439599..fafdba9 100644 --- a/dataset/expert_dataset.py +++ b/dataset/expert_dataset.py @@ -244,6 +244,7 @@ class ExpertTrajectoryDataset(Dataset): print(f" 观测维度: {obs_dim} (应为107)") if save_path: + os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: pickle.dump({ "trajectories": all_trajectories, @@ -282,7 +283,7 @@ if __name__ == "__main__": trajectories, observations = ExpertTrajectoryDataset.collect_with_full_obs( env_config, num_scenarios=10, - save_path="./expert_trajectories_full.pkl" + save_path="data/trajectories/expert_trajectories_full.pkl" ) if len(trajectories) > 0: diff --git a/logs/20260131-192601/events.out.tfevents.1769858761.Hfkk.41584.0 b/logs/20260131-192601/events.out.tfevents.1769858761.Hfkk.41584.0 new file mode 100644 index 0000000000000000000000000000000000000000..228172d3bfee5228ebdb7ce39657d3163ce48cb2 GIT binary patch literal 5884 zcmZwKc{G*zAHZ=-wrlJzvS%BrK{8X{B*~oynPjP>MqL_LbFIIw8`0F1B??WpRQ6Mo zrI3^*x!=jiQZg;WMDv@9lnG^tekk)B&**v1=RD8-^Pabl*XKFz=eW*d|GwtZo%7^g zfAGE%&Fyv6R@g1@=Z9H-8OaM1uCej7WGxr+!Xm;2yTknj0hST_LPPxp(UwsHzK|!d zvyP2gYr?h>cR$(czWQjV`KZNaTecpa=zg1B* z8P<2jcwzeGf+#0B52|JWvCPk+ss6Wp0Q7ogWD5b3pnT|LCmR50oa!OyS4BxYU^5g~+@|5m;NPrg6K(m3eSjV~28Gz(v$R2_QymWA+o~k_n*-j zgM?Q%|M;;p(A!2;!|s!p}-lLI8**PEKb4r3WRWFf`E$ zK;LqX$Pl1UXduap+1N!7xqv5R_a)gwP_OYB%vW&R4nVXKsI*!Jq^SS^isOyme z0aBoWO#Ackt0^*35}0Zsso_>pUEDAV%@VbUqpr!{|P`YcZw|tkO~d-OY;?M($|&%K(-xZ4?*tlD{$)U z$@2ijr)mb!lw%wUw>CZmK(dVu)&xkE2Kqw27UxxWeE^{6R1ZNvcwWc#kM6wyAhw8{ z&Hx%4JA$qZSbqUPPb=50B0y?1&>G)Hyk25-5P&|dB6|pWYjq7bpuBScw3DhCKrTNW zMgykpi@|AsJiYNFi1D;b{ZgO-HZWZ8XHRUPPTL@J+YNksAMd|lgEU97ha0$-*5SG{ zuA9IHx2c+8gO|03&{iGmNB}ZOnsOpQY#PYnXcHdw53vQHIjV=CO3n@JmQyhbKvu`d z=?tI;%f3f5ql@?ebg(sd69Lkof!0ml!NYCe@&U-MjqD*PX61eShsn-V01BjP29VQ* zXk_5Dc@TiO+H6k(q)7ucIo`+K+}g_kWTr#*5LC;#i<1X@3IWK6su@5QpGKk82}v9P z(o!n)AwXI*PzI+Pmj<_A0H88uvWK8}wO$-_@IWpAO;a@kD564$mX;Um0?^op++74{ zF%5Ll>@jwo{7wRZ8t2F!fC^xK-OlI?AVAtQkiqv) z@Xk{C5CFQFN%j!brrVEYw5BQn=p9uvfX06fN1CfdNdRQN9|sd49U5qSbsy$;s+0kc zcr@8V(1Y$L*h%C!2SBz|%>a@P3PnvN<}v{ES)84K0O`^|TdM|eXR#XmUh|&nA?V`B z2>vtwMj8NF#FNt*Kn|>6)VY^q20*sW`}Y$dJsN1D_7yhRF~kB-XoBh?NH1|1UwWh7 z0YHYg$mtBAnKHxPZbZI?IzDL8s35Fi5@s4R92H!gk~2p;Jjs)wK- zyJv85Re?MJ%~CZ3Xo*q)I<41U13+id!DIqtNCTPj-eRrlm>d8KSVQ&@q+&CTFT|!) z0Z=+sGl1M${ZY+VW$yr}WS>Pk0a`)>g&m#3i?*9?1fb+dvWK81&J1?`>pJ`wu#u`6 zK${b|DCS`Q3IO65y*f>RjA$VC@)=y^X{irD2F7F$L534Ec)gPg{8(aAHf|@pTomteu|vV06MGs4Vp!)#Tciu4oYHuNX2zuc=iwkPhIpA;T9aS@cTAul!q0$mn06LbJQ9*!A zX`mCW)7Yyq;~D_@X)4vkh5iZdLzfX3F!Z)+aSsS literal 0 HcmV?d00001 diff --git a/logs/20260131-194343/events.out.tfevents.1769859823.Hfkk.44173.0 b/logs/20260131-194343/events.out.tfevents.1769859823.Hfkk.44173.0 new file mode 100644 index 0000000000000000000000000000000000000000..29fd8bc7978e7578af7993485d6a4238d05636a3 GIT binary patch literal 1532 zcmZwGeMl2=9LMolPQ43B%}tTaR0th%!b;(i(g?#s!KAXFS|{CN( zveaA!{b8+8Ydx@PJGCs)d`P8HniO*|pL#N4=B(Jc?)Y+dfBty?KEKyLmm^yK{?z4| zka1{3#r+bl1@&=ei1(Uwp}B?NU8)1FkT391|`qf|PQ!bo4G zP^yYa#qx>l;+N~|{JZbRBjiYLCic>|D_etxqiAl4R-3BxJs>$al9efPd&Fz~bx&t~ zEMm#>*Wu`DG=!;3iDF6LRP`2QL!i^b*!vNI%dwzrgt>A$kSj+7>aRxa6U0Gy%2+~g8vdgDC3jnB(s|7$x zQv*J6(_9Qd=H3J23}__}^zG1~-8nIJ0)QI&*dBu3>Syh>4o3_C4REyp=%Mr~E{PsK z4nQ4?*dzlAcBk#x>1= z#5~aX)&x4pCLRKyxLUS{pjXZ`dg;`H9)NapwE*aE?z-PuX<=kh`%CKi&Q=4}f~Q5*HXyFc0+Tb0OWS%qIcJ z_=4>rsM4aKyLSwE0H}kj1whRSwfK_Swg-R)C32C$_uGf?Ktz|CP8mg`0H{sM_7HTo zzk*J3#7+Z{i>n1dVQtm8uJ=<70KHC=1~VXp2ihB0K{tD?bpX_w&h`+r@%uSCeK_eJ S06Do@0F;qgg|{dSqyGTl$nkps literal 0 HcmV?d00001 diff --git a/logs/20260131-200211/events.out.tfevents.1769860931.Hfkk.47961.0 b/logs/20260131-200211/events.out.tfevents.1769860931.Hfkk.47961.0 new file mode 100644 index 0000000000000000000000000000000000000000..93541ff97dc69e4f37eafce46896913a9d9d7dd6 GIT binary patch literal 227 zcmeZZfPjCKJmzv9egE=8b^I+yDc+=_#LPTB*Rs^S5-X!1JuaP+)V$*SqNM!9q7=R2 z(%js{qDsB;qRf)iBE3|Qs`#|boYZ)T$OX@ts&Y_sZ{86y;=Hth>n6xtEnzM}E-s(^ z;$r<0kOiW;;V$-@7oD0c166(E;8a;oiA_+|f?WJu9AMRDi8-QzY+r289|~Bg2vy5i n{pK5I5KOHYmk1ZHPikUOUS?i;d{JUas_3e1vn>Nn(k21`{_jms literal 0 HcmV?d00001 diff --git a/logs/20260131-202131/events.out.tfevents.1769862091.Hfkk.51799.0 b/logs/20260131-202131/events.out.tfevents.1769862091.Hfkk.51799.0 new file mode 100644 index 0000000000000000000000000000000000000000..82f1c3716b4a1f11f342c8408da319ff5499296b GIT binary patch literal 88 zcmeZZfPjCKJmzwq-}L!YP5doKDc+=_#LPTB*Rs^S5-X!1JuaP+)V$*SqNM!9q7=R2 h(%js{qDsB;qRf)iBE3|Qs`#|boYZ)T$cfk!TL6RLB0c~B literal 0 HcmV?d00001 diff --git a/logs/20260131-202136/events.out.tfevents.1769862096.Hfkk.51875.0 b/logs/20260131-202136/events.out.tfevents.1769862096.Hfkk.51875.0 new file mode 100644 index 0000000000000000000000000000000000000000..27481dc1f9b4396e72d667c7b428e4e095fd8e16 GIT binary patch literal 1581 zcmZwHZ%7ki9Kdm>mAfbZyD5lPT3G6sWF?&au@pl}OpGv6T9eLL=|;DIi5d7tQ&Z99 zG}qF=3eB8o4GG)L7n1!Y1z+Sph(ZO^LL{@5<<9GlA9v5ohwt-yKCgc19vf$^A>E3- zY~HneT`_(v%r~D@s?_o{1EJO>C1lHyOh>4-8Zuv_BrnRf`l2EwSu8h_Djh+}33sI^ zP!$r2jhld+XXm#ub3aZ7C1CL}7MFW&Y%mH$C_|&w#_QZ4NF|ji$tl(5vT&vz9&C)j zd}5h85cwl9T^p2z(rY#4hx5ja=OUR_RPyx)mioU9K@#MfK`2SJN_|N|Ds_ak({91d z*+a4&%!8PPjx5Ym?PEfl&>~k0fC{Y5 zc)^*)B>=*^SN&;_7Z3DhuHP}`-5CQw;WoC1Ad&2~qoZ^t4uCSbS^)GkqzM<*PY?jq zShE~TgEsL%C6-l(GooArK!#ejhajsfnEF#u=?6d$xmo};Z)(JI&~QHhjZ7qLr$L)} zpsTk6skXRBEdbOt$@UN=Bczm?yX=&$59-o9it15n>eZ!`_s!UHW7$5Ch1 z16ly8`^okYw^V>dVywkhiZhkp_u*p#D=AsI`+`5dbv8^$-*~prA}x zl_vms__5OkKmoQoJnF=<6@YRtj~<~x7!M>G*HdlNp<@6Pd4=sE=y!V=b)z-Q1VC9_ rEdV;0T#L`X`LGXc($7`x$qYz>g3#7%!nnTc;|cfb^;ka#?(5M%a~AA# literal 0 HcmV?d00001 diff --git a/logs/20260131-202739/events.out.tfevents.1769862459.Hfkk.53251.0 b/logs/20260131-202739/events.out.tfevents.1769862459.Hfkk.53251.0 new file mode 100644 index 0000000000000000000000000000000000000000..5a473fd3ec730a9310ab9d539912e40e1277740e GIT binary patch literal 1532 zcmZwFeMl2=9LMoyTHb|hGiM+^q-JQ!Emsj&KBQ7Dk)d<-PY>gCw>r|Pw>gVg58y*e zk2cXtDfEv>&{kSmv|S;DIeMT6>_O4cqVgfch*3Vo&UMF^yZiHa-{<$@$ElYrKJ)eJ z7UWHwTVvCI%m@oApb5%^ZMTvp%NCUuLz66|$y`Dgln}H5GnW?^6ZBE6f~G7ajgj6; zZ={My{oAm+SG@X&WDb86{3mzCB5kgAhg7XdZLT*ODiJW>FX_2}wYP@n+hSWhrjQ>{t>ddl|@v{Oqk z`3aGu0Q85iML@D!=kYkZ`Z)k4x7J=@L5Ki!rn-WW&$OEXXqN9G=w$gRCh7j$c>qdm Qk^z-bkiqx1{i44Xi}$^Q6_6bT#gz)Y^6*VGi@-FbiUehsicIYE7aw*(Mr*3%3J9Q zjK!30<7Rng?16&ab$hnQVLkFc&%L{cf^iv+XP7OPqgL++^2dghTjuv8J zS8gHzZO!`omH}-RfQB{5r5#^a0Z>3T*F#YDmv=6GbPEMQTD}$m^@ZLg^qmp60La83drZ0j`Ij+1g;_u5NY!&`Z7+0sXwwL_CZ_7XT>g!}Uo9%0JOr_BB1ttw+U(e*=Ydk{w$qiK>h+yU}q7E_@q<TGb6LG9XL<3RHde=NsA(hF+K1rlpxz#09Oo)`EJRwdI87Ecj7bhgf#!3~3#mNep zl2nMv#Y%EO7DGzjelq2krCqH%cj%#&3(sES;JbM4dTYXxAOhkN5}cKbKOnqt$J1j% zsgAJI)jgf1t9Z8ibhRa6L6}l?k~BtGD$8+GyWAOY(6#z0W(YR*4R!Wj5REzOZiM`(R$;SVzDjg|E-T{4^*3xiIxi2%rZ%rF+O3E*Mwu z@(=cN)Scc`o9xX?VQLN#@nHp;Xn!CApn?6OFbZVA0_~N5qVYL)))9cJNxFw1>#%ao zg3X#y0P1CG4$y(>)98g~p$34~ia!XaK!z+()Yw(cW3&+Eg`f3P4)-a47{cVu2#ujd0q^4SfJq0-zqI<^cUu zT#9OrnYaQF|2A05rt(5Y#Xz!ftJ6+yKbLjy|0Ov>+%!eDkLW zfSP_wi=#j$ED+ymE1r51s0ToK*XbUDR!{q2l~ipGK&?#80s6aKg_@^Y)ByCbB|M1& zEn|U%+7P@ksPQfU)wj|;1bve#!@eWV?f`U;sX0JZYfd8f2;oHl@(77cr9h@EkXmyP zcX^Lr1)wdVbPqw#LrDC{tMxhn#V|Doh~Ibu)ef~f0#M)Ok1{C`j|F;dnTpk#lGgy# z)j{_VbXS#)xAB8X0GefL4$wA_B9wonatPd{lZI|fAOT@Tn1zzbZ|-_$a`Ecbd-FW& zZ|ejLq8nBhPzwmwg6LNncp?0e8CdYNk?!Gw;hIF;)xTjoSYX>k*BlE31IN)V!{BfL zTIFL^Oo7Z;Ac0*z-aGe20|3o1Jp|>I6k}0y+Z6z^_oYwg0I8x2kVlZH8-UEFohm8N zJ1kIZUNP3ccQ^xp9A@Yqf{d1*#+!?N{}O=wnVJJsXPl1=IxD&WsN((9a}>y&1-e#W zh10KzDgo$#J>5f)XH*^DSN1#@fJ&K~0~B^H4-q%qZvv2G^zLx@7u|WOtmvF4AUJO7MpVK`A6@2z9c1^T{e||fdngcX7m4RlC z%%23HBRSh|Qy^;=X!m`LSG?4m2cTWKbPqv!gck44yj23;P$pAzfc(Gx3Z3##Bf-^< zM?PKxabE35y_~Ia0kuHDS`hnJC!Rz93c!NE|L7hrh-m1?mwSqO!GaW~=2$S= zo{l0%^?m^$y^Y;}QXnA+_YmZL{|261(IW()W~Sx<^*g1ao^_kD F{s(8?8#Dj_ literal 0 HcmV?d00001 diff --git a/logs/20260131-205450/events.out.tfevents.1769864090.Hfkk.59264.0 b/logs/20260131-205450/events.out.tfevents.1769864090.Hfkk.59264.0 new file mode 100644 index 0000000000000000000000000000000000000000..3cccc12fc41b79a1609888636149d992727bf2a0 GIT binary patch literal 88 zcmeZZfPjCKJmzw`J)QfkF8-FI6mL>dVrHJ6YguYuiIq{19+yr@YF=@EQBrM*}QKepaQD#YMkzOiDReV}zPHH?vWZy*2UjR;DAvpj5 literal 0 HcmV?d00001 diff --git a/logs/20260201-000321/events.out.tfevents.1769875401.Hfkk.64137.0 b/logs/20260201-000321/events.out.tfevents.1769875401.Hfkk.64137.0 new file mode 100644 index 0000000000000000000000000000000000000000..5a256d3da1f7fe6588bc528f8eab66e179734a76 GIT binary patch literal 15072 zcmZwOd0bBE|HttbOUjzERU}#|O46lLREkoGLfSVKib5J|!h|FVMTXIWD2YLq&P7?X zharr~l6@+&HG|*G`S!ccE^2NzdzqIU3YZom+$?4O;GLV zKAl42L!x8NoEJpI&KqgrZKgVSUPNp{TzqI;NPM_i!u;8@L*kRn662%iMZ}v${8t4< zM$d`}`rqFa#47srZC&^NLv@P|dEc_W{h)gd#$?&Y9j;X-uqsWs`TXofap8cBWS^r&ceM% z@oE5Sl4=g<&8E#{Wbf~z0F?2|&wWHtTN!9TP%oizhu1j(#qDH0f?VoMgp2_f#{(!& zsyUz^8d+qbYR5bP+3u~-6+!J}piE_R!QN+N2LKJ+$9e=E9`7QU*lE53&?KqmfL`Wg zl1ej&3IJta>1il}+RH%CmW&gkzD1@3Xzf+jBgoHhg0QFQ#uxwzQq2M7>TM)tCrBuO zyq!mxiXask=v+alps15&450pFS&yK^=?TK~$}@uiY z9b};W3JV1{uM0H*QuxYx1huV86?AsrFa*$0spf!e?9<8Va3cXu(q4tPEL%Yus_Lq0 z-Vurap7lc_{yTaV{{6GM(z!iMV3eI>BVM2?n;@rowXpAXyD>09v()1RPr7Cbw{{)Q zhY8Gb*y@}KG-}dF>76Z&0P;9_b+iallYu^v4MN{ldx`*La*Fi`nyir{xEtrI0%)34 zb3mCf>qy>FkL3W`{W{!L1a*{wtcPwDs#o@&2cYF|SdXAJRoTK#lL&7B9h7PgsHp8) z(tQ1qI)H)_k9&%sPBM^LkDUVDRCF9b_6u2$pcx-`3hFa%B>`xORC7R+^Ha(FDYG;I zG`dp9R|IvIf$kmKBP1Wc5CI^+D%K-t+oB?2=XKR&0Bw|N4ruk*pGn^9q8tF}ED!J( zL0x2^&s$1`g7xu#0jOT;5mf7NQkbl<+yOw{lG*AUkoDVD`-KD!o(pzbozlS?-Qjh(j_0O+ICBgjnS zoZxGCA8#@H`?A$Jpm7DuNv50m-*B`a$&5>E1#ypd4M1*E%>fx3CX)5}O8NkbQz+gcf_lk7N{!Ej3%dk(vpza3SZu5}wg)l$s?tt*Wpmv29A14sLnLzR`SAnwsVv)u0z zOkgy7$Yt?@zOo4%o!io7+kNl`$2gYtIKiF4YP95=(<7K5MyfdzJn)YsRa<)C4bEhP z$QvT4pA2-jUWLwCzr8T4VyCV||`J6+wD3(BVXN>inezpAgN| zVm*Ss_3K789<})mK-;C71L`?0fYdqU_5#r7%B0sKXn+jVXsSuQ|M-L-rpwQ<9zk9T z+Ef_!Y$kwA&$F5X+NM6841ay_6M(iw;ZF**- zQW=11q?!XdUpJKq$@j+r$T_R@lL*q6f$om#P4!~W?E_G!&8$bzrUP1ZRzur+0J4#4 z4k-PoFBy?mMFFIr?D|av8OT6Y5nA-H&4vsB)k{5si~l>KxFsjLC#t z9_0w2Yg;3g*Zuck*iZ(lchsVf-@d;EpzLhcBgn0bHvP?{vnGHpNHqskKF6D!2>en5 zM|)|yS({c6_h|3j_F)Q4Fsx8pRlLARHbLI2zO=r&Gd`AVkb0aTCb&Ocb9n-8sqP-O zI%k3(K3?SH^?APoXop9qt|G`-2CB=|r8^(2iG(f9@MJxL3MU!T^p^pJu%%T}%>nIo z@E|2WDYyYBbj^275o96*r8(%)xIuf*0LW`C>k$-Ts7o#L6!EcSs#J48g+tv*PGfB% zfc7|y)e%91WT0+$2hcmGO(y{;Y&7c;r1)!h`gYm%p|C?4Qq2Lq)gMp3mul()XugBr zKoK-p2Ab)wPs`inHv!0WH0u$Rmi3pA5H?^SfC8kN1Dd7lO4J<>zXDL7{k0|{$W#W( zJ7_?sTgJ@*P=nMXXhY**sx$6fAb@lau+=%B-nwH+`?XKD0q9Rt4|5S@CIh8?Go))@ zjlT|{95dDmU9bRx5?mv05ok*_5qM9>f!D0Zk)Le)QEP;>JS5{*omN_GSJgqCiLmTrc3~tonbwK44em1 zE$H2XZ6g62AEz3~Y21LFO`0q`wJGoc1{hKud#IkD!%j4e9n~b58)3 zNHqtvdg3Tj9^h05M|=KU>eLG29_=YBZl8xOHS&HuUcA6UHo@3!rc~2Hp({*KFZDRV z)3V|84fyb06MCkG*JW%lYwmW%xGx9N_@AB?#_AyHO|nd zV|FY&2%uc4=72u`Jd&82?8P6;r~X%`h#*TDC}!CZdOt60H|)^M0M;X@xxWqlTYGR1 z0QC-JH3t-5WK9xRzc>UtMCy(PilE^#(Dq~t`raiq8bBYV9zi#i2tCvPpap;opRm#OVo8-&X6*n_yECjukhSFqy1ll{2|%Nyngcrc&62#a z@6--JHJ^^o6hR|opq|E7v{>tn1%TFlW<7%bcx6TFW{30wP`OldK-0{Ik@Qol>i~4U z+CE+cS<661DkEst&F0wv%DBRM1UWphr;j!6zX6aS)f`Ykm^sPal+ztRnbCd=MG%pJ zZgeH|iMrMm00qXd9zoZaInm~nDbE0uDb*a%&pUk;(&h!c(f(PuS)+MQ-K2lTk_VA5~F?rJ#N|0o)~sujdN+Rq#=J_8dR>M*8;tJLy z=!2>)4cuO@0U%SU=74tk>XV!qKh^-Ktet(K2y&2tu9iDf;o6@k0hG|5^$6Ol?n|@Q zUta>CU!|G@ntpHq@i_7BBY+H!R~!&Q@~?a~HM`KI)?x;uJhG!-NC}&G;nFtyq+o6Y+Zq)9(az6lV-^zLfEe*G)9bc;8OI9t_9MDg* zb;--(e0;RMC2W4V2y&ExE*4Fo`>#aq0Z>>t>k;%h%!gW;eZB)bv{|Y-po_QqlE|Mb z@fq6XZC$Ada*}~G#(C14Iwo2G^2lX9g8qv1rpK;rR|inKRC7QUwmRf&fs+xuWak4@ z&Wj*t8OX$LBCWd8co9HLgIJHCC(nYZ;)h;)08}j19MFM%T0}K`$!$2=KZf{TXa#YP zcAG2JyJ3PQK?>Kz3&zSOXnW6__I1(y1{3%OvmPg?UE@c`fA_$*?=@1*nZQ)551H-s zGyy7KVmkc$j7Zlez!>2qr_fVOO4J%Wy>gwTf{PAr5it&(aEs3@Zs8M&+aE`XXZ zeW?{et}>ASPG73<;A=5}$}Y1WL2fjdURddfzj3OQY7VHkYEN?4-{=c~No@hZ>}s1Df-*CUJOX@CHDkZKk~tLF471+^O`| zeQ6}5%IUjG0Q6F-IiTY8>cq+^XE=cLqb4_tAa@xk!YY6!uSj_ZJM>oS5wu8s zCbgTMGao=&(QI`NXnBur#B0<)_~zH(THYdpCdfeRmj%*(3nK9ONBKC`BdF#64Ei`p zWdwj;Ni_%5J+BK<`@^gRcIedsC&je?9&LNbKp!oFY14=Wp0GoQidm1KP{(j;qHena zKrf`40}2?{nOOP#69}MVJxkk*AWs=+dv-9b*UQBFm4aTZN6^Aa<7su2?QsCzl4=fU z@`aAXBxTSYINHk%l2lqj+@pQMufNy91kVp1=`3F0C7Yn-WC)EouTuwGdaji9IDyua z7@GQSrVUKcBGsG;8UxkH=r(jDfHH>f(-c7yWuWzCA(SlGZ~{P~R;)+Rv(yA?{`z

5KM0!}|k`{ohrJ4h>wP{au zhrhH1(8VMUamSR5WP@-RNQJ0iX`StVd9Wmp`Q!8@>a`N~$>^ueEJR)nT>H0D5pylZYT+ z8K|&2iZ(l$-vf}_W!59er6P_xg)2D&C`GC{AmtItM8RT!CxDJPnT{4gQ)HlS4l#7- zu4cSFNN{F7f|mIu&})e7Ps?Mc7=if`Wro1RY*FYuF1prbdN zIwj7*-xf4UJx*|7SrolHvg#B}pr6H7=S*<1$3NB%?U(j|E&cnUTd)Y4E(2xFolTVk z3m(F2eCiSF5oBc*O;0;C<^t%oRC7Ra8^2m#|ET{PfUX5SpCN+$WuUR!Vri_=Gi3m+ z3uZln%9bpr&z^tu1<-M+=783!f42VSIs6KMZrqw5CxQZGpdkb2&^A?cFMxK`upU9f zj22NVy~Fqond?%`0jX?jww{$xwE;j{5vhqHC{PC4nm315d$!C4&}FGdkZGYlJ!G3T z3wG#}RC7SJhVQNad^-tm0ezfKEE7RNGSDXTIdq=oL1zRxvmQaos`F^zEW51$(iqEX z4yfT|gZ0r+75uLV_v{o?MNqH|v_B%ATHGkl1<=M(tVfXJ@5||cEzW-e=!{fzKqGwL zS$mXZN5D%qcin^x5fmZ=W!omuHRl!ap`K+L>k;(+K{8#}xm*oEVN%Tj?f&%IdYPdI ze$c8YoVQg3g~~v4%Mz%r|IMDTLreFt9zj7H7Sm5FI^PG-0jcJIba%Y8_78JXh8^1D zzVa6l6ea^*y*!U@QW$OnpxF~xkD$E9h16$t$zK5aMXEWVym9r`u}9+8!_nT$amvnC N5cg;|S-sTj{{Zi%FhBqR literal 0 HcmV?d00001 diff --git a/logs/bc/20260202-005850/events.out.tfevents.1769965130.Hfkk.240848.0 b/logs/bc/20260202-005850/events.out.tfevents.1769965130.Hfkk.240848.0 new file mode 100644 index 0000000000000000000000000000000000000000..0604e3d74735da814e90bedbfcae66d5806574e7 GIT binary patch literal 3080 zcmZwIeM}Q)9Kdlv1bU9qRv=Pbg>5h}QZ}|M9H79I4IG&pGNHOzT30sSifb#Z(Gq-v z@-oWmkQWitq`s69tO0T0>;r)WsfW#rSA-w?w1SJXeWaoF)^*YXNQMLKUC zlkQ7GY}Mjxw$)81ayPPOshRrp0F)O&_7HSK(qkUn z9b*qbMyh52HG5X#ojHH}06?3yjpqrFBMsD$v}7*m=x+m{)N-%D=rEb#;B-3_uxF%>ZJb)#8$=*aZNx8@fA6fLJt8*^UC_ zUsMPGL)WMtg6927P`SmS4uBkn$<-M^;gMx{=f7ReV3P8c*T!rhcZ`dzNRZ|}pY_2~ z>*%$6eyNjhMS>l?qv1~x4`8$%)XbM7L)3v%aHN4nWDj>Rd8iUS7}ZXI9qgoPh8?(H zFU6B0_QK!z>;5PA3D61}$g`yiJ-YI`5`aGamFyuXj#rDOGv68jphBu<0O`_8@cye? z+QA8FyhG**&`UIsYDkX)oTHCI&}y=WAYPjR@pJ`20Ca$=89?>Ri}6ae4;=0KgwCe~ z$dv}_%{hgneV#G^+Od`FA*e0Ogtj&2!EcR*su@6u1`XbSM}GsHkjc$rf7|*sewha9 z+Sh+@>vP8esPs>aMS$FBpt9g5bjH}! z4nXBIWDh|@&Of1hBQbse)Irq@pis+UJaU)nM*u1}dvgd7n+9qyw4&Z~8aTKvBC>~| zro79j`({-u0L@V~11P295I$LTP6I%8TT{FU5Qhd5twU%$Yv~hkLXC-J4?%22J96}& zkOR;is%8K!L?6T#z5WgYps>n8UjoFXfi#=D(TwTB1^`-HMfMOh6Lu0kj4J#cfVNXL z188n;KfbLtXcmB4c?$sqXeAAF#nOi)tkX^al;=(M5EO#MKpTj0v{x6$wf_fB1>XAr literal 0 HcmV?d00001 diff --git a/scripts/README.md b/scripts/README.md new file mode 100644 index 0000000..86de988 --- /dev/null +++ b/scripts/README.md @@ -0,0 +1,82 @@ +# 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_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_trained_policy.py) + +使用训练好的 **BC** 或 **MAGAIL** 模型在 45 维场景环境中运行,并实时渲染俯瞰图(top-down view)。统一入口:`scripts/visualize_trained_policy.py`。 + +**BC 模型**: +```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 +``` + +**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 +``` + +**自动推断类型**(根据 `--model_path` 扩展名:`.pt` → BC,否则 → MAGAIL): +```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 +``` + +**根目录 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`)。 + +--- + +### 数据分析与检查 + +| 脚本 | 用途 | 用法示例 | +|------|------|----------| +| [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 接口可能不一致,可选使用 | + +--- + +### 其他 + +| 脚本 | 用途 | 用法示例 | +|------|------|----------| +| [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. **可视化**:`visualize_trained_policy.py`(或根目录 `visualize_bc.py` 仅 BC)→ 从 `models/bc` 或 `models/magail` 加载模型,数据目录默认 `data/exp_filtered` diff --git a/scripts/README_visualize.md b/scripts/README_visualize.md deleted file mode 100644 index 76abadb..0000000 --- a/scripts/README_visualize.md +++ /dev/null @@ -1,113 +0,0 @@ -# 模型可视化脚本使用说明 - -## 功能 -使用训练好的MAGAIL模型在环境中运行,并生成俯瞰效果图(top-down view)。 - -## 使用方法 - -### 基本用法 - -```bash -python scripts/visualize_trained_model.py \ - --model_dir runs/magail_0113 \ - --episode 1250 \ - --data_dir data/exp_filtered \ - --num_scenarios 1 \ - --output_dir visualizations -``` - -### 参数说明 - -- `--model_dir`: 模型保存目录(例如:`runs/magail_0113`) -- `--episode`: 要加载的episode编号(例如:`1250`) -- `--data_dir`: Waymo数据目录(默认:`data/exp_filtered`) -- `--start_index`: 起始场景索引(默认:`0`) -- `--num_scenarios`: 要运行的场景数量(默认:`1`) -- `--horizon`: 每个episode的最大步数(默认:`200`) -- `--output_dir`: 输出图像保存目录(默认:`visualizations`) -- `--save_all_frames`: 保存所有帧(否则按间隔保存) -- `--save_interval`: 保存帧的间隔,当不使用`--save_all_frames`时生效(默认:`10`) -- `--gif_duration`: GIF每帧持续时间(毫秒),默认50ms(20fps)。值越小,GIF播放越快 - -### 示例 - -#### 1. 查看最新训练的模型(episode 1250) -```bash -python scripts/visualize_trained_model.py \ - --model_dir runs/magail_0113 \ - --episode 1250 \ - --num_scenarios 3 \ - --output_dir visualizations/episode_1250 -``` - -#### 2. 保存所有帧(用于制作视频) -```bash -python scripts/visualize_trained_model.py \ - --model_dir runs/magail_0113 \ - --episode 1250 \ - --save_all_frames \ - --output_dir visualizations/episode_1250_all_frames -``` - -#### 3. 每5步保存一帧 -```bash -python scripts/visualize_trained_model.py \ - --model_dir runs/magail_0113 \ - --episode 1250 \ - --save_interval 5 \ - --output_dir visualizations/episode_1250_sparse -``` - -#### 4. 生成更快的GIF(30fps) -```bash -python scripts/visualize_trained_model.py \ - --model_dir runs/magail_0113 \ - --episode 1250 \ - --gif_duration 33 \ - --output_dir visualizations/episode_1250 -``` - -## 输出 - -脚本会在指定的输出目录中创建以下文件: -- `scenario_{idx}.gif`: **场景动画GIF**(主要输出) -- `scenario_{idx}_step_{step:04d}.png`: 每个保存步骤的俯瞰图(可选) -- `scenario_{idx}_final.png`: 每个场景的最终状态图 - -### GIF格式 -- 分辨率:1600x900 -- 格式:GIF动画 -- 包含完整的场景运行过程 -- 显示场景编号、步数、智能体数量和奖励信息 -- 默认帧率:20fps(可通过`--gif_duration`调整) - -### 图像格式 -- 分辨率:1600x900 -- 格式:PNG -- 包含语义地图和车辆轨迹 - -## 注意事项 - -1. **GPU要求**: 脚本需要CUDA支持,如果没有GPU会自动使用CPU(速度较慢) -2. **渲染模式**: 使用MetaDrive的top-down渲染模式,会弹出窗口显示实时渲染 -3. **内存占用**: 如果保存所有帧,会占用较多磁盘空间 -4. **场景数据**: 确保`--data_dir`指向正确的Waymo数据目录 - -## 故障排除 - -### 模型文件不存在 -``` -FileNotFoundError: 模型文件不存在: runs/magail_0113/model_1250_actor.pth -``` -**解决**: 检查模型目录和episode编号是否正确 - -### 场景数据不存在 -``` -ValueError: Data directory not found -``` -**解决**: 确保`--data_dir`指向正确的数据目录 - -### 渲染失败 -如果遇到渲染相关错误,可以尝试: -- 降低`film_size`参数(在脚本中修改) -- 使用无头模式(需要修改脚本) diff --git a/scripts/generate_expert_data.py b/scripts/generate_expert_data.py index d08e137..853300f 100644 --- a/scripts/generate_expert_data.py +++ b/scripts/generate_expert_data.py @@ -153,8 +153,8 @@ def generate_data(args): if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument("--data_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/exp_filtered", help="Path to Waymo pickles (or filtered index)") - parser.add_argument("--output_dir", type=str, default="/home/huangfukk/MAGAIL4AutoDrive/data/training", help="Output directory") + parser.add_argument("--data_dir", type=str, default="data/exp_filtered", help="Path to Waymo pickles (or filtered index)") + parser.add_argument("--output_dir", type=str, default="data/training_data", help="Output directory") parser.add_argument("--start_index", type=int, default=0) parser.add_argument("--num_scenarios", type=int, default=10) diff --git a/scripts/visualize_trained_policy.py b/scripts/visualize_trained_policy.py index 3919b64..83e1bee 100644 --- a/scripts/visualize_trained_policy.py +++ b/scripts/visualize_trained_policy.py @@ -1,143 +1,189 @@ +""" +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 -import time # 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 train_magail import Actor, MAGAILScenarioEnv +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): - # 1. Load Environment - data_path = os.path.abspath(args.data_dir) + 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, # Visualisation enabled + "use_render": True, "sequential_seed": True, "start_scenario_index": args.start_index, "num_scenarios": args.num_scenarios, "log_level": 40, } - - print("Initializing MAGAILScenarioEnv...") + + print(f"Initializing BCScenarioEnv (policy_type={policy_type})...") try: - env = MAGAILScenarioEnv(config=env_config, agent2policy={}) + env = BCScenarioEnv(env_config, agent2policy={}) except Exception as e: print(f"Error init env: {e}. Trying to close lingering engine...") try: close_engine() - except: + except Exception: pass - env = MAGAILScenarioEnv(config=env_config, agent2policy={}) + env = BCScenarioEnv(env_config, agent2policy={}) - # 2. Load Model state_dim = 45 action_dim = 2 - - actor = Actor(state_dim, action_dim).cuda() - - model_path = args.model_path - if not os.path.exists(model_path): - # Try to find it in runs/ - potential_path = os.path.join("runs", "magail_production", model_path) - if os.path.exists(potential_path): - model_path = potential_path - else: - # Try appending _actor.pth - potential_path = model_path + "_actor.pth" - if os.path.exists(potential_path): - model_path = potential_path - else: - raise ValueError(f"Model path {args.model_path} not found.") - + 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}...") - actor.load_state_dict(torch.load(model_path)) - actor.eval() - - # 3. Run Loop + + 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} ---") - - # Reset try: - # Use sequential seed logic or specific seed? - # ExpertReplayEnv/ScenarioEnv logic: seed matches scenario index if configured right obs_dict = env.reset(seed=i) except Exception as e: print(f"Error resetting {i}: {e}. Skipping.") - # Try soft reset try: close_engine() - env = MAGAILScenarioEnv(config=env_config, agent2policy={}) - except: + 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 = {} - # Inference - for agent_id, obs in obs_dict.items(): - # Preprocess obs: (45,) -> (1, 45) tensor - obs_tensor = torch.FloatTensor(obs).unsqueeze(0).cuda() - with torch.no_grad(): + 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) - # Deterministic action for viz? Or sample? - # Usually deterministic (mean) is better for checking performance - # But training uses sample. if args.deterministic: - action = torch.tanh(dist.mean) # Use mean of Gaussian + actions_np = torch.tanh(dist.mean).cpu().numpy() else: - pre_tanh = dist.sample() - action = torch.tanh(pre_tanh) - - actions[agent_id] = action.cpu().numpy().flatten() - - # Step + 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) - - # Render + episode_reward += sum(rewards.values()) + env.render( mode="top_down", text={ "Scenario": i, "Step": step_count, - "Agents": len(obs_dict) - } + "Agents": len(obs_dict), + "Total Reward": f"{episode_reward:.2f}", + }, ) - step_count += 1 - # time.sleep(0.02) # Slow down if needed - + if dones["__all__"] or step_count >= args.horizon: - print(f"Scenario finished at step {step_count}") + 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() - parser.add_argument("--model_path", type=str, required=True, help="Path to actor model pth (e.g. runs/magail_production/model_50_actor.pth)") - parser.add_argument("--data_dir", type=str, default="data/exp_filtered") + 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="Use mean action instead of sampling") - + parser.add_argument("--deterministic", action="store_true", help="For MAGAIL: use mean action instead of sampling") + args = parser.parse_args() visualize_model(args) diff --git a/train_bc.py b/train_bc.py new file mode 100644 index 0000000..74ad5f8 --- /dev/null +++ b/train_bc.py @@ -0,0 +1,186 @@ +""" +BC 训练脚本:负责数据加载、环境评估、日志与保存;BC 算法由 Algorithm.bc 提供。 +使用方式不变:python train_bc.py [--expert_data_path data/training_data] [--save_dir models/bc] ... +""" +import os +import glob +import pickle +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 + + +def load_expert_data(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 + + +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_data(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) diff --git a/train_magail.py b/train_magail.py index ec2eebd..061027e 100644 --- a/train_magail.py +++ b/train_magail.py @@ -10,6 +10,7 @@ import signal import sys from torch.utils.data import DataLoader from dataset.magail_dataset import MAGAILExpertDataset +from Env.bc_env import BCScenarioEnv # --- Networks --- @@ -161,58 +162,10 @@ class PPO: torch.save(self.actor.state_dict(), checkpoint_path + "_actor.pth") torch.save(self.critic.state_dict(), checkpoint_path + "_critic.pth") -from Env.scenario_env import MultiAgentScenarioEnv - -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 - # --- Training Loop --- def train(args): - # 1. Setup Environment (Dummy for now, usually you run simulation here) - # But for MAGAIL we need to collect generated trajectories. - # We need the Env class to be importable. - from Env.scenario_env import MultiAgentScenarioEnv - from Env.simple_idm_policy import ConstantVelocityPolicy # Just for init - + # 1. Setup Environment (45-dim obs via BCScenarioEnv) # Config for Env env_config = { "data_directory": args.data_dir, @@ -255,12 +208,7 @@ def train(args): yield batch expert_iter = cycle(expert_loader) - # 4. Initialize Env - from Env.expert_replay_env import ExpertReplayEnv # Using ReplayEnv for config, but we need ScenarioEnv for simulation? - # Actually we need MultiAgentScenarioEnv for interactive training, not Replay. - from Env.scenario_env import MultiAgentScenarioEnv - from Env.simple_idm_policy import ConstantVelocityPolicy # Placeholder policy for init - + # 4. Initialize Env (BCScenarioEnv provides 45-dim obs) # 2. Setup Models # Determine state dim from environment if possible, or use fixed # Expert data has 45 dim? @@ -316,7 +264,7 @@ def train(args): # obs_dict[agent_id] = obs # return obs_dict - env = MAGAILScenarioEnv(config=env_config, agent2policy={}) # Pass empty dict if we control all externally + env = BCScenarioEnv(env_config, agent2policy={}) # 45-dim obs print("Starting training...") @@ -401,7 +349,7 @@ def train(args): import gc gc.collect() - env = MAGAILScenarioEnv(config=env_config, agent2policy={}) + env = BCScenarioEnv(env_config, agent2policy={}) obs_dict = env.reset(seed=seed) episode_reward = 0 @@ -575,7 +523,7 @@ def train(args): print(f"Episode {i_episode}: Disc Loss {disc_loss.item():.4f} | PPO Loss {ppo_loss:.4f} | Mean Reward {np.mean(all_gail_rewards):.4f}") if i_episode % 50 == 0: - ppo_agent.save(os.path.join(args.log_dir, f"model_{i_episode}")) + ppo_agent.save(os.path.join(args.save_dir, f"model_{i_episode}")) env.close() if writer: @@ -588,11 +536,13 @@ if __name__ == '__main__': parser.add_argument("--batch_size", type=int, default=1024) parser.add_argument("--max_episodes", type=int, default=1000) parser.add_argument("--num_scenarios", type=int, default=100) - parser.add_argument("--log_dir", type=str, default="runs/magail_exp") + parser.add_argument("--log_dir", type=str, default="logs/magail", help="TensorBoard log directory") + parser.add_argument("--save_dir", type=str, default="models/magail", help="Directory to save model checkpoints") args = parser.parse_args() - # Create log dir + # Create log dir and save dir os.makedirs(args.log_dir, exist_ok=True) + os.makedirs(args.save_dir, exist_ok=True) train(args) diff --git a/visualize_bc.py b/visualize_bc.py new file mode 100644 index 0000000..2ac0946 --- /dev/null +++ b/visualize_bc.py @@ -0,0 +1,17 @@ +""" +Thin wrapper: forwards to scripts/visualize_trained_policy.py --policy_type bc. +Use: python visualize_bc.py [--model_path models/bc/policy_best.pt] [other args...] +Or call directly: python scripts/visualize_trained_policy.py --policy_type bc --model_path models/bc/policy_best.pt +""" +import subprocess +import sys +import os + +def main(): + script_dir = os.path.dirname(os.path.abspath(__file__)) + script = os.path.join(script_dir, "scripts", "visualize_trained_policy.py") + cmd = [sys.executable, script, "--policy_type", "bc"] + sys.argv[1:] + sys.exit(subprocess.run(cmd).returncode) + +if __name__ == "__main__": + main()