diff --git a/Env/__pycache__/expert_replay_env.cpython-313.pyc b/Env/__pycache__/expert_replay_env.cpython-313.pyc index 5feb9b1..a3f53e1 100644 Binary files a/Env/__pycache__/expert_replay_env.cpython-313.pyc and b/Env/__pycache__/expert_replay_env.cpython-313.pyc differ diff --git a/Env/__pycache__/expert_replay_env.cpython-39.pyc b/Env/__pycache__/expert_replay_env.cpython-39.pyc index b0935b9..fcb7475 100644 Binary files a/Env/__pycache__/expert_replay_env.cpython-39.pyc and b/Env/__pycache__/expert_replay_env.cpython-39.pyc differ diff --git a/Env/__pycache__/inverse_dynamics.cpython-39.pyc b/Env/__pycache__/inverse_dynamics.cpython-39.pyc index fbac54f..b41c57c 100644 Binary files a/Env/__pycache__/inverse_dynamics.cpython-39.pyc and b/Env/__pycache__/inverse_dynamics.cpython-39.pyc differ diff --git a/Env/__pycache__/scenario_env.cpython-39.pyc b/Env/__pycache__/scenario_env.cpython-39.pyc index d924eb8..918470d 100644 Binary files a/Env/__pycache__/scenario_env.cpython-39.pyc and b/Env/__pycache__/scenario_env.cpython-39.pyc differ diff --git a/Env/bc_env.py b/Env/bc_env.py index b802fba..59c080c 100644 --- a/Env/bc_env.py +++ b/Env/bc_env.py @@ -1,13 +1,113 @@ from Env.scenario_env import MultiAgentScenarioEnv +from Env.utils import filter_traffic_tracks_to_birth_lists +from metadrive.component.vehicle.vehicle_type import DefaultVehicle import numpy as np + class BCScenarioEnv(MultiAgentScenarioEnv): """ Environment for Behavior Cloning Evaluation. Uses the same 45-dim observation as ExpertReplayEnv: - Ego State (5): x, y, vx, vy, heading - Neighbors (40): 10 nearest * (rel_x, rel_y, vx, vy) + + Spawns background (static) vehicles so that observation distribution matches expert data collection: + expert data is generated with ExpertReplayEnv which includes bg_* in active_agents, so the policy + was trained on obs that can include those neighbors. Demo should use the same scene for consistency. """ + def reset(self, seed=None): + # Clear background vehicles from previous episode so engine.reset() passes _object_clean_check + if getattr(self, "engine", None) is not None: + ids_bg = [ + oid for oid, obj in self.engine.get_objects().items() + if (getattr(obj, "name", None) or getattr(obj, "id", None) or "").startswith("bg_") + ] + if ids_bg: + self.engine.clear_objects(ids_bg, force_destroy=True) + for aid in list(self.engine.agent_manager.active_agents.keys()): + if aid.startswith("bg_"): + self.engine.agent_manager.active_agents.pop(aid, None) + obs = super().reset(seed=seed) + self._spawn_background_vehicles() + return self._get_all_obs() + + def _build_birth_lists_from_traffic(self): + """Same lane/static filter as expert data; return background_vehicles so we spawn them (match training obs).""" + car_birth_info_list, background_vehicles, obj_to_clean, stats = filter_traffic_tracks_to_birth_lists( + self.engine.traffic_manager.current_traffic_data, + self.engine.traffic_manager.sdc_scenario_id, + self.engine.map_manager, + return_stats=True, + ) + if stats["n_controlled"] == 0 and stats["n_total"] > 0: + print( + "[BCScenarioEnv] 0 controlled agents: total_vehicles={}, off_lane={}, static={}, no_valid={}.".format( + stats["n_total"], + stats["n_off_lane"], + stats["n_static"], + stats["n_no_valid"], + ) + ) + return car_birth_info_list, background_vehicles, obj_to_clean + + def _spawn_background_vehicles(self): + """Spawn static background vehicles so they appear in active_agents and thus in obs (same as ExpertReplayEnv).""" + for sid, car in self.background_vehicles.items(): + if car["show_time"] != self.round: + continue + bg_id = f"bg_{car['id']}" + if bg_id in self.engine.agent_manager.active_agents: + continue + vehicle_config = {} + if "length" in car and "width" in car: + vehicle_config = {"length": car["length"], "width": car["width"]} + v = self.engine.spawn_object( + DefaultVehicle, + name=bg_id, + vehicle_config=vehicle_config, + position=car["begin"], + heading=car["heading"], + ) + v.set_velocity([0, 0]) + self.engine.agent_manager.active_agents[bg_id] = v + v.valid_mask = car["valid"] + v.start_t = car["show_time"] + + def _update_background_vehicles(self): + self._spawn_background_vehicles() + to_remove = [] + objects_to_clear = [] + for aid, v in self.engine.agent_manager.active_agents.items(): + if not aid.startswith("bg_"): + continue + if hasattr(v, "valid_mask"): + if self.round >= len(v.valid_mask) or not v.valid_mask[self.round]: + 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([v.id for v in objects_to_clear]) + + def step(self, action_dict): + self.round += 1 + for agent_id, action in action_dict.items(): + if agent_id in self.controlled_agents: + self.controlled_agents[agent_id].before_step(action) + self.engine.step() + self.engine.after_step() + for agent_id in action_dict: + if agent_id in self.controlled_agents: + self.controlled_agents[agent_id].after_step() + self._spawn_controlled_agents() + self._update_background_vehicles() + 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.episode_step >= self.config["horizon"] + infos = {aid: {} 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 = {} diff --git a/Env/expert_replay_env.py b/Env/expert_replay_env.py index ce516f0..d850244 100644 --- a/Env/expert_replay_env.py +++ b/Env/expert_replay_env.py @@ -35,132 +35,44 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): 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.background_vehicles = {} self.expert_tracks = {} - # Capture SDC track for ego replay (MetaDrive default agent) self.sdc_track = None self.sdc_vehicle = None + + # 在加载新场景前,必须清除上一轮通过 spawn_object 生成的物体,否则 engine.reset() 内 _object_clean_check 会报错 + # 从 engine 当前对象中按名称筛选(与 manager 无关的对象需在此清理),并强制销毁 + ids_to_clear = [] + for oid, obj in self.engine.get_objects().items(): + name = getattr(obj, "name", None) or getattr(obj, "id", None) + if name and (str(name).startswith("controlled_") or str(name).startswith("bg_")): + ids_to_clear.append(oid) + if ids_to_clear: + self.engine.clear_objects(ids_to_clear, force_destroy=True) + self.controlled_agents.clear() + self.controlled_agent_ids.clear() + for aid in list(self.engine.agent_manager.active_agents.keys()): + if aid.startswith("bg_") or aid.startswith("controlled_"): + self.engine.agent_manager.active_agents.pop(aid, None) + if self.replay_sdc and hasattr(self.engine, "traffic_manager"): sdc_sid = self.engine.traffic_manager.sdc_scenario_id self.sdc_track = self.engine.traffic_manager.current_traffic_data.get(sdc_sid, None) - _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: + from Env.utils import filter_traffic_tracks_to_birth_lists + traffic_data = self.engine.traffic_manager.current_traffic_data + car_birth_info_list, self.background_vehicles, obj_to_clean = filter_traffic_tracks_to_birth_lists( + traffic_data, + self.engine.traffic_manager.sdc_scenario_id, + self.engine.map_manager, + ) + for entry in car_birth_info_list: + sid = entry["scenario_id"] + if sid in traffic_data: + self.expert_tracks[sid] = traffic_data[sid] + self.car_birth_info_list = car_birth_info_list + for scenario_id in obj_to_clean: self.engine.traffic_manager.current_traffic_data.pop(scenario_id) - # --- MODIFIED SECTION END --- self.engine.reset() self.reset_sensors() @@ -260,38 +172,6 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): 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 @@ -314,7 +194,7 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): self.engine.agent_manager.active_agents.pop(aid, None) if objects_to_clear: - self.engine.clear_objects(objects_to_clear) + self.engine.clear_objects([v.id for v in objects_to_clear]) def _spawn_controlled_agents(self): for car in self.car_birth_info_list: diff --git a/Env/scenario_env.py b/Env/scenario_env.py index 8488854..6f1f531 100644 --- a/Env/scenario_env.py +++ b/Env/scenario_env.py @@ -76,28 +76,9 @@ class MultiAgentScenarioEnv(ScenarioEnv): if self.engine is None: raise ValueError("Broken MetaDrive instance.") - # 记录专家数据中每辆车的位置,接着全部清除,只保留位置等信息,用于后续生成 - _obj_to_clean_this_frame = [] - self.car_birth_info_list = [] - for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items(): - if scenario_id == self.engine.traffic_manager.sdc_scenario_id: - continue - else: - if track["type"] == MetaDriveType.VEHICLE: - _obj_to_clean_this_frame.append(scenario_id) - valid = track['state']['valid'] - first_show = np.argmax(valid) if valid.any() else -1 - last_show = len(valid) - 1 - np.argmax(valid[::-1]) if valid.any() else -1 - # id,出现时间,出生点坐标,出生朝向,目的地 - self.car_birth_info_list.append({ - 'id': track['metadata']['object_id'], - 'show_time': first_show, - 'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]), - 'heading': track['state']['heading'][first_show], - 'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]) - }) - - for scenario_id in _obj_to_clean_this_frame: + self.background_vehicles = getattr(self, "background_vehicles", {}) + self.car_birth_info_list, self.background_vehicles, _obj_to_clean = self._build_birth_lists_from_traffic() + for scenario_id in _obj_to_clean: self.engine.traffic_manager.current_traffic_data.pop(scenario_id) # Clear vehicles we spawned via engine.spawn_object() so _object_clean_check() passes @@ -126,6 +107,27 @@ class MultiAgentScenarioEnv(ScenarioEnv): return self._get_all_obs() + def _build_birth_lists_from_traffic(self): + """Build car_birth_info_list and obj_to_clean from current_traffic_data. Override for filtered (lane/static) selection.""" + _obj_to_clean_this_frame = [] + car_birth_info_list = [] + for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items(): + if scenario_id == self.engine.traffic_manager.sdc_scenario_id: + continue + if track["type"] == MetaDriveType.VEHICLE: + _obj_to_clean_this_frame.append(scenario_id) + valid = track["state"]["valid"] + first_show = int(np.argmax(valid)) if valid.any() else -1 + last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else -1 + car_birth_info_list.append({ + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": (track["state"]["position"][first_show, 0], track["state"]["position"][first_show, 1]), + "heading": track["state"]["heading"][first_show], + "end": (track["state"]["position"][last_show, 0], track["state"]["position"][last_show, 1]), + }) + return car_birth_info_list, {}, _obj_to_clean_this_frame + def _spawn_controlled_agents(self): # ego_vehicle = self.engine.agent_manager.active_agents.get("default_agent") # ego_position = ego_vehicle.position if ego_vehicle else np.array([0, 0]) diff --git a/Env/utils.py b/Env/utils.py index c19bf24..8dfced6 100644 --- a/Env/utils.py +++ b/Env/utils.py @@ -2,6 +2,143 @@ import numpy as np import torch import random +from metadrive.type import MetaDriveType + + +def is_on_lane(pos, map_manager, threshold=2.0): + """Check if a position is on a valid lane (within lateral tolerance).""" + if map_manager is None or map_manager.current_map is None: + return True + try: + lane, _ = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True) + if lane is None: + return False + long, lat = lane.local_coordinates(pos) + width = lane.width + if abs(lat) <= (width / 2 + threshold): + return True + return False + except Exception: + return False + + +def filter_traffic_tracks_to_birth_lists( + current_traffic_data, + sdc_scenario_id, + map_manager, + *, + lane_threshold=5.0, + static_displacement_threshold=5.0, + static_speed_threshold=1.0, + return_stats=False, +): + """ + Filter traffic tracks into controlled (car_birth_info_list) and background lists. + + - controlled (car_birth_info_list): 非 SDC、类型 VEHICLE、至少一帧 valid、在车道内、且非静态 + (位移/速度超过阈值)。用于策略控制或专家回放,spawn 时机为 show_time == round。 + - background (background_vehicles): 同上但在车道内且判定为静态(位移 < 5m、速度 < 1 m/s)。 + 仅作场景占位与观测邻居,spawn 时机为 show_time == round,按 valid 在 step 中移除。 + + Returns (car_birth_info_list, background_vehicles, obj_to_clean) or, if return_stats=True, + (car_birth_info_list, background_vehicles, obj_to_clean, stats_dict). + stats_dict: n_total, n_no_valid, n_off_lane, n_static, n_controlled. + """ + car_birth_info_list = [] + background_vehicles = {} + obj_to_clean = [] + n_total = 0 + n_no_valid = 0 + n_off_lane = 0 + n_static = 0 + + for scenario_id, track in current_traffic_data.items(): + if scenario_id == sdc_scenario_id: + continue + if track["type"] != MetaDriveType.VEHICLE: + continue + + n_total += 1 + obj_to_clean.append(scenario_id) + valid = track["state"]["valid"] + if not valid.any(): + n_no_valid += 1 + continue + + first_show = int(np.argmax(valid)) + last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) + mid_show = (first_show + last_show) // 2 + + start_pos = track["state"]["position"][first_show] + is_valid_track = True + if not is_on_lane(start_pos, map_manager, threshold=lane_threshold): + mid_pos = track["state"]["position"][mid_show] + if not is_on_lane(mid_pos, map_manager, threshold=lane_threshold): + is_valid_track = False + + if not is_valid_track: + n_off_lane += 1 + continue + + positions = track["state"]["position"][valid.astype(bool)] + velocities = track["state"]["velocity"][valid.astype(bool)] + total_displacement = 0.0 + max_speed = 0.0 + if len(positions) > 1: + total_displacement = float(np.linalg.norm(positions[-1] - positions[0])) + max_speed = float(np.max(np.linalg.norm(velocities, axis=1))) + is_static = total_displacement < static_displacement_threshold and max_speed < static_speed_threshold + + if is_static: + n_static += 1 + background_vehicles[scenario_id] = { + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": ( + float(track["state"]["position"][first_show, 0]), + float(track["state"]["position"][first_show, 1]), + ), + "heading": float(track["state"]["heading"][first_show]), + "end": ( + float(track["state"]["position"][last_show, 0]), + float(track["state"]["position"][last_show, 1]), + ), + "scenario_id": scenario_id, + "length": track["state"]["length"][first_show], + "width": track["state"]["width"][first_show], + "valid": valid, + } + continue + + car_birth_info_list.append({ + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": ( + float(track["state"]["position"][first_show, 0]), + float(track["state"]["position"][first_show, 1]), + ), + "heading": float(track["state"]["heading"][first_show]), + "end": ( + float(track["state"]["position"][last_show, 0]), + float(track["state"]["position"][last_show, 1]), + ), + "scenario_id": scenario_id, + "length": track["state"]["length"][first_show], + "width": track["state"]["width"][first_show], + }) + + if return_stats: + stats = { + "n_total": n_total, + "n_no_valid": n_no_valid, + "n_off_lane": n_off_lane, + "n_static": n_static, + "n_controlled": len(car_birth_info_list), + } + return car_birth_info_list, background_vehicles, obj_to_clean, stats + return car_birth_info_list, background_vehicles, obj_to_clean + + def set_seed(seed): if seed == -1: seed = np.random.randint(0, 10000) diff --git a/scripts/README.md b/scripts/README.md index 0dcffa1..349bd1f 100644 --- a/scripts/README.md +++ b/scripts/README.md @@ -35,7 +35,7 @@ python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500 ``` -- **policy**(BC 或 MAGAIL 训练策略): +- **policy**(BC 或 MAGAIL 训练策略):与专家数据生成/回放一致——同一套车道+静态筛选、且会生成背景车(bg_*),使观测分布与训练集一致,便于在训练集上公平演示。 ```bash python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1 python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth --num_scenarios 1 --deterministic diff --git a/scripts/visualize.py b/scripts/visualize.py index 7b411db..5a3e532 100644 --- a/scripts/visualize.py +++ b/scripts/visualize.py @@ -177,6 +177,9 @@ def _run_policy(args): continue print(f"Scenario loaded. Controlled agents: {len(obs_dict)}") + if len(obs_dict) == 0: + print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.") + continue step_count = 0 episode_reward = 0.0