Compare commits
3 Commits
03dee0205a
...
BC_dev
| Author | SHA1 | Date | |
|---|---|---|---|
| 0f9f080e77 | |||
| ceb6648a31 | |||
| 95cc78d940 |
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,13 +1,99 @@
|
|||||||
from Env.scenario_env import MultiAgentScenarioEnv
|
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
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
class BCScenarioEnv(MultiAgentScenarioEnv):
|
class BCScenarioEnv(MultiAgentScenarioEnv):
|
||||||
"""
|
"""
|
||||||
Environment for Behavior Cloning Evaluation.
|
Environment for Behavior Cloning Evaluation.
|
||||||
Uses the same 45-dim observation as ExpertReplayEnv:
|
Uses the same 45-dim observation as ExpertReplayEnv:
|
||||||
- Ego State (5): x, y, vx, vy, heading
|
- Ego State (5): x, y, vx, vy, heading
|
||||||
- Neighbors (40): 10 nearest * (rel_x, rel_y, vx, vy)
|
- 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 all static background vehicles once at reset (no show_time filter; same as ExpertReplayEnv)."""
|
||||||
|
for sid, car in self.background_vehicles.items():
|
||||||
|
bg_id = f"bg_{car['id']}"
|
||||||
|
if bg_id in self.engine.agent_manager.active_agents:
|
||||||
|
continue
|
||||||
|
vehicle_config = {}
|
||||||
|
if "length" in car and "width" in car:
|
||||||
|
vehicle_config = {"length": car["length"], "width": car["width"]}
|
||||||
|
v = self.engine.spawn_object(
|
||||||
|
DefaultVehicle,
|
||||||
|
name=bg_id,
|
||||||
|
vehicle_config=vehicle_config,
|
||||||
|
position=car["begin"],
|
||||||
|
heading=car["heading"],
|
||||||
|
)
|
||||||
|
v.set_velocity([0, 0])
|
||||||
|
self.engine.agent_manager.active_agents[bg_id] = v
|
||||||
|
v.valid_mask = car.get("valid")
|
||||||
|
v.start_t = car.get("show_time")
|
||||||
|
|
||||||
|
def _update_background_vehicles(self):
|
||||||
|
# Static vehicles are spawned once at init and never removed.
|
||||||
|
pass
|
||||||
|
|
||||||
|
def step(self, action_dict):
|
||||||
|
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):
|
def _get_all_obs(self):
|
||||||
# Implement custom observation: 30m range, 10 nearest vehicles
|
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||||
obs_dict = {}
|
obs_dict = {}
|
||||||
|
|||||||
@@ -35,132 +35,44 @@ class ExpertReplayEnv(MultiAgentScenarioEnv):
|
|||||||
if self.engine is None:
|
if self.engine is None:
|
||||||
raise ValueError("Broken MetaDrive instance.")
|
raise ValueError("Broken MetaDrive instance.")
|
||||||
|
|
||||||
self.background_vehicles = {} # Vehicles that exist but are static/background
|
self.background_vehicles = {}
|
||||||
|
|
||||||
# Helper function to check if a position is on a valid lane
|
|
||||||
def is_on_lane(pos, map_manager, threshold=2.0):
|
|
||||||
# Check if point is close to any lane in the road network
|
|
||||||
# This can be expensive if checked for every point, so we check sample points
|
|
||||||
# or rely on lane index if available.
|
|
||||||
# Waymo tracks don't have lane index, just positions.
|
|
||||||
# We can use map.road_network.get_closest_lane_index(pos)
|
|
||||||
if map_manager is None or map_manager.current_map is None:
|
|
||||||
return True # If no map, assume valid
|
|
||||||
|
|
||||||
try:
|
|
||||||
# Use a larger search radius to catch slightly offset lanes
|
|
||||||
lane, lane_index = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True)
|
|
||||||
if lane is None:
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check lateral distance
|
|
||||||
long, lat = lane.local_coordinates(pos)
|
|
||||||
width = lane.width
|
|
||||||
# Allow being slightly off-lane (e.g. changing lanes)
|
|
||||||
# But parking lots are usually far from defined lanes in Waymo converted maps
|
|
||||||
if abs(lat) <= (width / 2 + threshold):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
# --- MODIFIED SECTION START ---
|
|
||||||
# Capture expert tracks before they are cleaned
|
|
||||||
self.expert_tracks = {}
|
self.expert_tracks = {}
|
||||||
# Capture SDC track for ego replay (MetaDrive default agent)
|
|
||||||
self.sdc_track = None
|
self.sdc_track = None
|
||||||
self.sdc_vehicle = 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"):
|
if self.replay_sdc and hasattr(self.engine, "traffic_manager"):
|
||||||
sdc_sid = self.engine.traffic_manager.sdc_scenario_id
|
sdc_sid = self.engine.traffic_manager.sdc_scenario_id
|
||||||
self.sdc_track = self.engine.traffic_manager.current_traffic_data.get(sdc_sid, None)
|
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
|
from Env.utils import filter_traffic_tracks_to_birth_lists
|
||||||
|
traffic_data = self.engine.traffic_manager.current_traffic_data
|
||||||
for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items():
|
car_birth_info_list, self.background_vehicles, obj_to_clean = filter_traffic_tracks_to_birth_lists(
|
||||||
if scenario_id == self.engine.traffic_manager.sdc_scenario_id:
|
traffic_data,
|
||||||
continue
|
self.engine.traffic_manager.sdc_scenario_id,
|
||||||
else:
|
self.engine.map_manager,
|
||||||
if track["type"] == MetaDriveType.VEHICLE:
|
)
|
||||||
_obj_to_clean_this_frame.append(scenario_id)
|
for entry in car_birth_info_list:
|
||||||
|
sid = entry["scenario_id"]
|
||||||
valid = track['state']['valid']
|
if sid in traffic_data:
|
||||||
if not valid.any():
|
self.expert_tracks[sid] = traffic_data[sid]
|
||||||
continue
|
self.car_birth_info_list = car_birth_info_list
|
||||||
|
for scenario_id in obj_to_clean:
|
||||||
first_show = np.argmax(valid)
|
|
||||||
last_show = len(valid) - 1 - np.argmax(valid[::-1])
|
|
||||||
mid_show = (first_show + last_show) // 2
|
|
||||||
|
|
||||||
# 1. Lane check (existing logic)
|
|
||||||
points_to_check = [first_show, mid_show, last_show]
|
|
||||||
on_road_count = 0
|
|
||||||
is_valid_track = True
|
|
||||||
start_pos = track['state']['position'][first_show]
|
|
||||||
if not is_on_lane(start_pos, self.engine.map_manager, threshold=5.0): # 5m tolerance
|
|
||||||
mid_pos = track['state']['position'][mid_show]
|
|
||||||
if not is_on_lane(mid_pos, self.engine.map_manager, threshold=5.0):
|
|
||||||
is_valid_track = False
|
|
||||||
|
|
||||||
# 2. Static check
|
|
||||||
# Calculate total displacement and max speed
|
|
||||||
positions = track['state']['position'][valid.astype(bool)]
|
|
||||||
velocities = track['state']['velocity'][valid.astype(bool)]
|
|
||||||
|
|
||||||
total_displacement = 0
|
|
||||||
max_speed = 0
|
|
||||||
if len(positions) > 1:
|
|
||||||
total_displacement = np.linalg.norm(positions[-1] - positions[0])
|
|
||||||
max_speed = np.max(np.linalg.norm(velocities, axis=1))
|
|
||||||
|
|
||||||
is_static = False
|
|
||||||
if total_displacement < 5.0 and max_speed < 1.0: # Relaxed threshold: <5m move and <1m/s
|
|
||||||
is_static = True
|
|
||||||
|
|
||||||
# Decision logic:
|
|
||||||
# - If off-road AND static: Skip completely (don't even spawn as background)
|
|
||||||
# - If off-road but moving: Maybe keep? Or skip? Usually off-road moving is weird, skip.
|
|
||||||
# - If on-road but static: Spawn as BACKGROUND (visible but not controlled agent)
|
|
||||||
# - If on-road and moving: Spawn as CONTROLLED agent
|
|
||||||
|
|
||||||
if not is_valid_track:
|
|
||||||
# Skip off-road vehicles entirely (both static and moving off-road)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if is_static:
|
|
||||||
# Add to background list, but NOT to car_birth_info_list (which is for controlled agents)
|
|
||||||
# We need a way to spawn them. Let's add a separate list.
|
|
||||||
self.background_vehicles[scenario_id] = {
|
|
||||||
'id': track['metadata']['object_id'],
|
|
||||||
'show_time': first_show,
|
|
||||||
'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]),
|
|
||||||
'heading': track['state']['heading'][first_show],
|
|
||||||
'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]),
|
|
||||||
'scenario_id': scenario_id,
|
|
||||||
'length': track['state']['length'][first_show],
|
|
||||||
'width': track['state']['width'][first_show],
|
|
||||||
'valid': valid # Need validity to know when to show/hide
|
|
||||||
}
|
|
||||||
continue # Do not add to controlled list
|
|
||||||
|
|
||||||
# Store the full track for replay (only for controlled agents)
|
|
||||||
self.expert_tracks[scenario_id] = track
|
|
||||||
|
|
||||||
self.car_birth_info_list.append({
|
|
||||||
'id': track['metadata']['object_id'],
|
|
||||||
'show_time': first_show,
|
|
||||||
'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]),
|
|
||||||
'heading': track['state']['heading'][first_show],
|
|
||||||
'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]),
|
|
||||||
'scenario_id': scenario_id, # Keep track of original ID to lookup tracks
|
|
||||||
'length': track['state']['length'][first_show],
|
|
||||||
'width': track['state']['width'][first_show]
|
|
||||||
})
|
|
||||||
|
|
||||||
for scenario_id in _obj_to_clean_this_frame:
|
|
||||||
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||||
# --- MODIFIED SECTION END ---
|
|
||||||
|
|
||||||
self.engine.reset()
|
self.engine.reset()
|
||||||
self.reset_sensors()
|
self.reset_sensors()
|
||||||
@@ -185,7 +97,7 @@ class ExpertReplayEnv(MultiAgentScenarioEnv):
|
|||||||
# We covered most of it.
|
# We covered most of it.
|
||||||
|
|
||||||
self._spawn_controlled_agents()
|
self._spawn_controlled_agents()
|
||||||
self._spawn_background_vehicles() # Initial spawn for background
|
self._spawn_all_background_vehicles_at_init()
|
||||||
|
|
||||||
# Ensure SDC/ego is moved to the correct initial expert state.
|
# Ensure SDC/ego is moved to the correct initial expert state.
|
||||||
if self.replay_sdc:
|
if self.replay_sdc:
|
||||||
@@ -202,119 +114,33 @@ class ExpertReplayEnv(MultiAgentScenarioEnv):
|
|||||||
|
|
||||||
return self._get_all_obs()
|
return self._get_all_obs()
|
||||||
|
|
||||||
def _spawn_background_vehicles(self):
|
def _spawn_all_background_vehicles_at_init(self):
|
||||||
# Spawn static/background vehicles
|
"""Spawn all static background vehicles once at reset (no show_time filter; no removal by valid)."""
|
||||||
# Since they are static, we might just spawn them once if their show_time is 0
|
|
||||||
# But Waymo tracks have valid bits, they might appear/disappear.
|
|
||||||
# For optimization, if they are truly static (never move), we just spawn them when show_time matches.
|
|
||||||
|
|
||||||
# We need to track spawned background vehicles to remove them if they become invalid?
|
|
||||||
# Since we defined them as "static", they probably stay put.
|
|
||||||
# But validity might change (e.g. late spawn).
|
|
||||||
|
|
||||||
# For simplicity in this step, let's just iterate and spawn if time matches
|
|
||||||
for sid, car in self.background_vehicles.items():
|
for sid, car in self.background_vehicles.items():
|
||||||
if car['show_time'] == self.round:
|
bg_id = f"bg_{car['id']}"
|
||||||
# Spawn as a Traffic Vehicle (not PolicyVehicle), or just a static object?
|
if bg_id in self.engine.agent_manager.active_agents:
|
||||||
# Using DefaultVehicle is fine, but don't add to controlled_agents
|
continue
|
||||||
|
vehicle_config = {}
|
||||||
# Check duplication
|
if 'length' in car and 'width' in car:
|
||||||
bg_id = f"bg_{car['id']}"
|
vehicle_config = {
|
||||||
# if bg_id in self.engine.obj_to_id: # obj_to_id might not be available in all versions
|
"length": car['length'],
|
||||||
if bg_id in self.engine.agent_manager.active_agents:
|
"width": car['width']
|
||||||
continue
|
}
|
||||||
|
v = self.engine.spawn_object(
|
||||||
vehicle_config = {}
|
DefaultVehicle,
|
||||||
if 'length' in car and 'width' in car:
|
name=bg_id,
|
||||||
vehicle_config = {
|
vehicle_config=vehicle_config,
|
||||||
"length": car['length'],
|
position=car['begin'],
|
||||||
"width": car['width']
|
heading=car['heading']
|
||||||
}
|
)
|
||||||
|
v.set_velocity([0, 0])
|
||||||
v = self.engine.spawn_object(
|
self.engine.agent_manager.active_agents[bg_id] = v
|
||||||
DefaultVehicle,
|
v.valid_mask = car.get('valid')
|
||||||
name=bg_id,
|
v.start_t = car['show_time']
|
||||||
vehicle_config=vehicle_config,
|
|
||||||
position=car['begin'],
|
|
||||||
heading=car['heading']
|
|
||||||
)
|
|
||||||
|
|
||||||
# Set color to grey/dark to indicate background
|
|
||||||
v.set_velocity([0, 0])
|
|
||||||
# Maybe set color? MetaDrive vehicles random color.
|
|
||||||
# v.set_color(...) if supported
|
|
||||||
|
|
||||||
# Register as an active object but NOT controlled agent
|
|
||||||
# The engine manages it.
|
|
||||||
# CRITICAL: We need it in self.engine.agent_manager.active_agents for Observation?
|
|
||||||
# If we want it to be seen by Lidar/Observation, it needs to be an "agent" or "traffic".
|
|
||||||
# DefaultVehicle spawned this way is just an object.
|
|
||||||
# We should add it to traffic manager? Or just leave it as object?
|
|
||||||
# MultiAgentScenarioEnv._get_all_obs iterates self.engine.agent_manager.active_agents
|
|
||||||
|
|
||||||
# If we want it in observation, we must add it to active_agents OR iterate over all objects.
|
|
||||||
# Adding to active_agents is easier for compatibility.
|
|
||||||
self.engine.agent_manager.active_agents[bg_id] = v
|
|
||||||
|
|
||||||
# Store valid mask to remove it later if needed?
|
|
||||||
v.valid_mask = car['valid']
|
|
||||||
v.start_t = car['show_time']
|
|
||||||
|
|
||||||
def _update_background_vehicles(self):
|
def _update_background_vehicles(self):
|
||||||
# Remove background vehicles if they become invalid
|
# Static vehicles are spawned once at init and never removed (no spawn/remove by show_time or valid).
|
||||||
# Or spawn new ones
|
pass
|
||||||
self._spawn_background_vehicles()
|
|
||||||
|
|
||||||
# Check validity for existing
|
|
||||||
to_remove = []
|
|
||||||
for aid, v in self.engine.agent_manager.active_agents.items():
|
|
||||||
if aid.startswith("bg_"):
|
|
||||||
# Check validity
|
|
||||||
if hasattr(v, 'valid_mask'):
|
|
||||||
curr_step = self.round
|
|
||||||
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
|
||||||
to_remove.append(aid)
|
|
||||||
|
|
||||||
for aid in to_remove:
|
|
||||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
|
||||||
# if aid in self.engine.obj_to_id:
|
|
||||||
# self.engine.clear_objects([self.engine.obj_to_id[aid]])
|
|
||||||
# Instead, we should find the object by ID and clear it.
|
|
||||||
# Since we don't track obj directly, we can't easily clear it without obj ref.
|
|
||||||
# Wait, active_agents stores the vehicle object.
|
|
||||||
# So we can just clear that object.
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Re-iterate to clear objects properly
|
|
||||||
for aid in to_remove:
|
|
||||||
# We need to find the vehicle object to clear it.
|
|
||||||
# But we popped it from active_agents.
|
|
||||||
# Wait, we should get it before pop.
|
|
||||||
pass
|
|
||||||
|
|
||||||
def _update_background_vehicles(self):
|
|
||||||
# Remove background vehicles if they become invalid
|
|
||||||
# Or spawn new ones
|
|
||||||
self._spawn_background_vehicles()
|
|
||||||
|
|
||||||
# Check validity for existing
|
|
||||||
to_remove = []
|
|
||||||
objects_to_clear = []
|
|
||||||
|
|
||||||
for aid, v in self.engine.agent_manager.active_agents.items():
|
|
||||||
if aid.startswith("bg_"):
|
|
||||||
# Check validity
|
|
||||||
if hasattr(v, 'valid_mask'):
|
|
||||||
curr_step = self.round
|
|
||||||
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
|
||||||
to_remove.append(aid)
|
|
||||||
objects_to_clear.append(v)
|
|
||||||
|
|
||||||
for aid in to_remove:
|
|
||||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
|
||||||
|
|
||||||
if objects_to_clear:
|
|
||||||
self.engine.clear_objects(objects_to_clear)
|
|
||||||
|
|
||||||
def _spawn_controlled_agents(self):
|
def _spawn_controlled_agents(self):
|
||||||
for car in self.car_birth_info_list:
|
for car in self.car_birth_info_list:
|
||||||
|
|||||||
@@ -64,6 +64,11 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
self.round = 0
|
self.round = 0
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_controlled_in_scenario(self) -> int:
|
||||||
|
"""整个场景中受控车轨迹总数(car_birth_info_list 长度),会在不同 show_time 陆续 spawn。"""
|
||||||
|
return len(getattr(self, "car_birth_info_list", []))
|
||||||
|
|
||||||
def reset(self, seed: Union[None, int] = None):
|
def reset(self, seed: Union[None, int] = None):
|
||||||
self.round = 0
|
self.round = 0
|
||||||
if self.logger is None:
|
if self.logger is None:
|
||||||
@@ -76,28 +81,14 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
if self.engine is None:
|
if self.engine is None:
|
||||||
raise ValueError("Broken MetaDrive instance.")
|
raise ValueError("Broken MetaDrive instance.")
|
||||||
|
|
||||||
# 记录专家数据中每辆车的位置,接着全部清除,只保留位置等信息,用于后续生成
|
# 注意:_build_birth_lists_from_traffic() 在 engine.reset() 之前执行,读的是当前 engine 的
|
||||||
_obj_to_clean_this_frame = []
|
# current_traffic_data 与 map_manager.current_map。若复用同一 env 连续 reset(0)、reset(1),
|
||||||
self.car_birth_info_list = []
|
# MetaDrive 可能已按 seed 更新了 traffic 为 scenario 1,但 map 仍为 scenario 0(在 engine.reset() 才切图),
|
||||||
for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items():
|
# 导致 is_on_lane( scenario_1 车位, scenario_0 地图 ) 全为 False → 全部 off_lane → 0 受控车。
|
||||||
if scenario_id == self.engine.traffic_manager.sdc_scenario_id:
|
# 因此多场景时应“每个 scenario 单独建 env”(start_scenario_index=i, num_scenarios=1)再 reset(seed=i)。
|
||||||
continue
|
self.background_vehicles = getattr(self, "background_vehicles", {})
|
||||||
else:
|
self.car_birth_info_list, self.background_vehicles, _obj_to_clean = self._build_birth_lists_from_traffic()
|
||||||
if track["type"] == MetaDriveType.VEHICLE:
|
for scenario_id in _obj_to_clean:
|
||||||
_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.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||||
|
|
||||||
# Clear vehicles we spawned via engine.spawn_object() so _object_clean_check() passes
|
# Clear vehicles we spawned via engine.spawn_object() so _object_clean_check() passes
|
||||||
@@ -126,6 +117,27 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
|
|
||||||
return self._get_all_obs()
|
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):
|
def _spawn_controlled_agents(self):
|
||||||
# ego_vehicle = self.engine.agent_manager.active_agents.get("default_agent")
|
# ego_vehicle = self.engine.agent_manager.active_agents.get("default_agent")
|
||||||
# ego_position = ego_vehicle.position if ego_vehicle else np.array([0, 0])
|
# ego_position = ego_vehicle.position if ego_vehicle else np.array([0, 0])
|
||||||
|
|||||||
221
Env/utils.py
221
Env/utils.py
@@ -2,6 +2,227 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
import random
|
import random
|
||||||
|
|
||||||
|
from metadrive.type import MetaDriveType
|
||||||
|
|
||||||
|
|
||||||
|
def _static_obbs_overlap(car_a, car_b):
|
||||||
|
"""
|
||||||
|
Check if two static vehicle OBBs overlap (2D SAT).
|
||||||
|
car_a, car_b: dicts with "begin" (x, y), "heading" (rad), "length", "width".
|
||||||
|
begin is center; half-extents are length/2, width/2.
|
||||||
|
"""
|
||||||
|
def _get_corners(car):
|
||||||
|
cx, cy = car["begin"][0], car["begin"][1]
|
||||||
|
h = float(car["heading"])
|
||||||
|
L2 = float(car["length"]) / 2.0
|
||||||
|
W2 = float(car["width"]) / 2.0
|
||||||
|
ux, uy = np.cos(h), np.sin(h)
|
||||||
|
vx, vy = -np.sin(h), np.cos(h)
|
||||||
|
return np.array([
|
||||||
|
[cx + L2 * ux + W2 * vx, cy + L2 * uy + W2 * vy],
|
||||||
|
[cx + L2 * ux - W2 * vx, cy + L2 * uy - W2 * vy],
|
||||||
|
[cx - L2 * ux - W2 * vx, cy - L2 * uy - W2 * vy],
|
||||||
|
[cx - L2 * ux + W2 * vx, cy - L2 * uy + W2 * vy],
|
||||||
|
])
|
||||||
|
|
||||||
|
def _get_axes(car):
|
||||||
|
h = float(car["heading"])
|
||||||
|
return [
|
||||||
|
np.array([np.cos(h), np.sin(h)]),
|
||||||
|
np.array([-np.sin(h), np.cos(h)]),
|
||||||
|
]
|
||||||
|
|
||||||
|
corners_a = _get_corners(car_a)
|
||||||
|
corners_b = _get_corners(car_b)
|
||||||
|
axes = _get_axes(car_a) + _get_axes(car_b)
|
||||||
|
|
||||||
|
for axis in axes:
|
||||||
|
proj_a = corners_a @ axis
|
||||||
|
proj_b = corners_b @ axis
|
||||||
|
min_a, max_a = proj_a.min(), proj_a.max()
|
||||||
|
min_b, max_b = proj_b.min(), proj_b.max()
|
||||||
|
if max_a < min_b or max_b < min_a:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _deduplicate_background_by_collision(background_vehicles):
|
||||||
|
"""
|
||||||
|
Merge static tracks that collide (same physical vehicle). Build collision graph,
|
||||||
|
find connected components, keep one representative per component (min show_time, then scenario_id).
|
||||||
|
"""
|
||||||
|
if not background_vehicles:
|
||||||
|
return background_vehicles
|
||||||
|
items = list(background_vehicles.items())
|
||||||
|
n = len(items)
|
||||||
|
# Build adjacency by index
|
||||||
|
parent = list(range(n))
|
||||||
|
|
||||||
|
def find(i):
|
||||||
|
if parent[i] != i:
|
||||||
|
parent[i] = find(parent[i])
|
||||||
|
return parent[i]
|
||||||
|
|
||||||
|
def union(i, j):
|
||||||
|
pi, pj = find(i), find(j)
|
||||||
|
if pi != pj:
|
||||||
|
parent[pi] = pj
|
||||||
|
|
||||||
|
for i in range(n):
|
||||||
|
for j in range(i + 1, n):
|
||||||
|
if _static_obbs_overlap(items[i][1], items[j][1]):
|
||||||
|
union(i, j)
|
||||||
|
|
||||||
|
# Representative per component: index with min (show_time, scenario_id)
|
||||||
|
comp_rep = {}
|
||||||
|
for i in range(n):
|
||||||
|
r = find(i)
|
||||||
|
sid, car = items[i][0], items[i][1]
|
||||||
|
key = (car.get("show_time", 0), sid)
|
||||||
|
if r not in comp_rep or key < comp_rep[r][0]:
|
||||||
|
comp_rep[r] = (key, sid, car)
|
||||||
|
|
||||||
|
return {sid: car for (_, sid, car) in comp_rep.values()}
|
||||||
|
|
||||||
|
|
||||||
|
def is_on_lane(pos, map_manager, threshold=2.0):
|
||||||
|
"""Check if a position is on a valid lane (within lateral tolerance)."""
|
||||||
|
if map_manager is None or map_manager.current_map is None:
|
||||||
|
return True
|
||||||
|
try:
|
||||||
|
lane, _ = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True)
|
||||||
|
if lane is None:
|
||||||
|
return False
|
||||||
|
long, lat = lane.local_coordinates(pos)
|
||||||
|
width = lane.width
|
||||||
|
if abs(lat) <= (width / 2 + threshold):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def filter_traffic_tracks_to_birth_lists(
|
||||||
|
current_traffic_data,
|
||||||
|
sdc_scenario_id,
|
||||||
|
map_manager,
|
||||||
|
*,
|
||||||
|
lane_threshold=5.0,
|
||||||
|
static_displacement_threshold=5.0,
|
||||||
|
static_speed_threshold=1.0,
|
||||||
|
return_stats=False,
|
||||||
|
deduplicate_static_by_collision=True,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Filter traffic tracks into controlled (car_birth_info_list) and background lists.
|
||||||
|
|
||||||
|
- controlled (car_birth_info_list): 非 SDC、类型 VEHICLE、至少一帧 valid、在车道内、且非静态
|
||||||
|
(位移/速度超过阈值)。用于策略控制或专家回放,spawn 时机为 show_time == round。
|
||||||
|
- background (background_vehicles): 同上但在车道内且判定为静态(位移 < 5m、速度 < 1 m/s)。
|
||||||
|
仅作场景占位与观测邻居,spawn 时机为 show_time == round,按 valid 在 step 中移除。
|
||||||
|
|
||||||
|
Returns (car_birth_info_list, background_vehicles, obj_to_clean) or, if return_stats=True,
|
||||||
|
(car_birth_info_list, background_vehicles, obj_to_clean, stats_dict).
|
||||||
|
stats_dict: n_total, n_no_valid, n_off_lane, n_static, n_controlled.
|
||||||
|
"""
|
||||||
|
car_birth_info_list = []
|
||||||
|
background_vehicles = {}
|
||||||
|
obj_to_clean = []
|
||||||
|
n_total = 0
|
||||||
|
n_no_valid = 0
|
||||||
|
n_off_lane = 0
|
||||||
|
n_static = 0
|
||||||
|
|
||||||
|
for scenario_id, track in current_traffic_data.items():
|
||||||
|
if scenario_id == sdc_scenario_id:
|
||||||
|
continue
|
||||||
|
if track["type"] != MetaDriveType.VEHICLE:
|
||||||
|
continue
|
||||||
|
|
||||||
|
n_total += 1
|
||||||
|
obj_to_clean.append(scenario_id)
|
||||||
|
valid = track["state"]["valid"]
|
||||||
|
if not valid.any():
|
||||||
|
n_no_valid += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
first_show = int(np.argmax(valid))
|
||||||
|
last_show = len(valid) - 1 - int(np.argmax(valid[::-1]))
|
||||||
|
mid_show = (first_show + last_show) // 2
|
||||||
|
|
||||||
|
start_pos = track["state"]["position"][first_show]
|
||||||
|
is_valid_track = True
|
||||||
|
if not is_on_lane(start_pos, map_manager, threshold=lane_threshold):
|
||||||
|
mid_pos = track["state"]["position"][mid_show]
|
||||||
|
if not is_on_lane(mid_pos, map_manager, threshold=lane_threshold):
|
||||||
|
is_valid_track = False
|
||||||
|
|
||||||
|
if not is_valid_track:
|
||||||
|
n_off_lane += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
positions = track["state"]["position"][valid.astype(bool)]
|
||||||
|
velocities = track["state"]["velocity"][valid.astype(bool)]
|
||||||
|
total_displacement = 0.0
|
||||||
|
max_speed = 0.0
|
||||||
|
if len(positions) > 1:
|
||||||
|
total_displacement = float(np.linalg.norm(positions[-1] - positions[0]))
|
||||||
|
max_speed = float(np.max(np.linalg.norm(velocities, axis=1)))
|
||||||
|
is_static = total_displacement < static_displacement_threshold and max_speed < static_speed_threshold
|
||||||
|
|
||||||
|
if is_static:
|
||||||
|
n_static += 1
|
||||||
|
background_vehicles[scenario_id] = {
|
||||||
|
"id": track["metadata"]["object_id"],
|
||||||
|
"show_time": first_show,
|
||||||
|
"begin": (
|
||||||
|
float(track["state"]["position"][first_show, 0]),
|
||||||
|
float(track["state"]["position"][first_show, 1]),
|
||||||
|
),
|
||||||
|
"heading": float(track["state"]["heading"][first_show]),
|
||||||
|
"end": (
|
||||||
|
float(track["state"]["position"][last_show, 0]),
|
||||||
|
float(track["state"]["position"][last_show, 1]),
|
||||||
|
),
|
||||||
|
"scenario_id": scenario_id,
|
||||||
|
"length": track["state"]["length"][first_show],
|
||||||
|
"width": track["state"]["width"][first_show],
|
||||||
|
"valid": valid,
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
|
||||||
|
car_birth_info_list.append({
|
||||||
|
"id": track["metadata"]["object_id"],
|
||||||
|
"show_time": first_show,
|
||||||
|
"begin": (
|
||||||
|
float(track["state"]["position"][first_show, 0]),
|
||||||
|
float(track["state"]["position"][first_show, 1]),
|
||||||
|
),
|
||||||
|
"heading": float(track["state"]["heading"][first_show]),
|
||||||
|
"end": (
|
||||||
|
float(track["state"]["position"][last_show, 0]),
|
||||||
|
float(track["state"]["position"][last_show, 1]),
|
||||||
|
),
|
||||||
|
"scenario_id": scenario_id,
|
||||||
|
"length": track["state"]["length"][first_show],
|
||||||
|
"width": track["state"]["width"][first_show],
|
||||||
|
})
|
||||||
|
|
||||||
|
if deduplicate_static_by_collision and background_vehicles:
|
||||||
|
background_vehicles = _deduplicate_background_by_collision(background_vehicles)
|
||||||
|
|
||||||
|
if return_stats:
|
||||||
|
stats = {
|
||||||
|
"n_total": n_total,
|
||||||
|
"n_no_valid": n_no_valid,
|
||||||
|
"n_off_lane": n_off_lane,
|
||||||
|
"n_static": n_static,
|
||||||
|
"n_controlled": len(car_birth_info_list),
|
||||||
|
}
|
||||||
|
return car_birth_info_list, background_vehicles, obj_to_clean, stats
|
||||||
|
return car_birth_info_list, background_vehicles, obj_to_clean
|
||||||
|
|
||||||
|
|
||||||
def set_seed(seed):
|
def set_seed(seed):
|
||||||
if seed == -1:
|
if seed == -1:
|
||||||
seed = np.random.randint(0, 10000)
|
seed = np.random.randint(0, 10000)
|
||||||
|
|||||||
@@ -49,6 +49,35 @@ def load_expert_pkl(expert_data_path):
|
|||||||
return obs_data, act_data
|
return obs_data, act_data
|
||||||
|
|
||||||
|
|
||||||
|
def get_expert_scenario_ids(expert_data_path, max_ids=10):
|
||||||
|
"""
|
||||||
|
从专家 pkl 中收集出现过的 scenario_id(这些场景在采集时曾有受控车)。
|
||||||
|
用于 eval 时只在这些场景上评估,保证 eval 有受控车。
|
||||||
|
返回排序后的 list,最多 max_ids 个。
|
||||||
|
"""
|
||||||
|
if os.path.isdir(expert_data_path):
|
||||||
|
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
||||||
|
elif os.path.exists(expert_data_path):
|
||||||
|
pkl_files = [expert_data_path]
|
||||||
|
else:
|
||||||
|
return []
|
||||||
|
|
||||||
|
seen = set()
|
||||||
|
for pkl_file in pkl_files:
|
||||||
|
try:
|
||||||
|
with open(pkl_file, "rb") as f:
|
||||||
|
data = pickle.load(f)
|
||||||
|
if isinstance(data, list):
|
||||||
|
for traj in data:
|
||||||
|
if "scenario_id" in traj:
|
||||||
|
seen.add(traj["scenario_id"])
|
||||||
|
# dict 格式通常没有 per-trajectory scenario_id,跳过
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
out = sorted(seen)[:max_ids]
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
class MAGAILExpertDataset(Dataset):
|
class MAGAILExpertDataset(Dataset):
|
||||||
def __init__(self, data_dir, transform=None):
|
def __init__(self, data_dir, transform=None):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -35,7 +35,7 @@
|
|||||||
python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500
|
python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500
|
||||||
```
|
```
|
||||||
|
|
||||||
- **policy**(BC 或 MAGAIL 训练策略):
|
- **policy**(BC 或 MAGAIL 训练策略):与专家数据生成/回放一致——同一套车道+静态筛选、且会生成背景车(bg_*),使观测分布与训练集一致,便于在训练集上公平演示。
|
||||||
```bash
|
```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 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
|
python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth --num_scenarios 1 --deterministic
|
||||||
|
|||||||
@@ -30,32 +30,34 @@ def _run_replay(args):
|
|||||||
max_available = len(summary_lookup) - args.start_index
|
max_available = len(summary_lookup) - args.start_index
|
||||||
num_to_run = min(args.num_scenarios, max_available)
|
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}...")
|
print(f"Initializing ExpertReplayEnv with data from {data_path}...")
|
||||||
env = ExpertReplayEnv(config=env_config)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
for i in range(args.start_index, args.start_index + num_to_run):
|
for i in range(args.start_index, args.start_index + num_to_run):
|
||||||
print(f"\n--- Playing Scenario {i} ---")
|
print(f"\n--- Playing Scenario {i} ---")
|
||||||
|
# Each scenario uses a fresh env (start_scenario_index=i, num_scenarios=1) so the second
|
||||||
|
# scenario and beyond are fully cleaned and loaded like a single-scenario run.
|
||||||
|
env_config = {
|
||||||
|
"data_directory": data_path,
|
||||||
|
"is_multi_agent": True,
|
||||||
|
"num_controlled_agents": 100,
|
||||||
|
"horizon": args.horizon,
|
||||||
|
"use_render": True,
|
||||||
|
"sequential_seed": True,
|
||||||
|
"reactive_traffic": False,
|
||||||
|
"start_scenario_index": i,
|
||||||
|
"num_scenarios": 1,
|
||||||
|
"log_level": 40,
|
||||||
|
}
|
||||||
|
env = ExpertReplayEnv(config=env_config)
|
||||||
try:
|
try:
|
||||||
obs = env.reset(seed=i)
|
obs = env.reset(seed=i)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error resetting scenario {i}: {e}")
|
print(f"Error resetting scenario {i}: {e}")
|
||||||
|
env.close()
|
||||||
continue
|
continue
|
||||||
|
|
||||||
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}")
|
print(f"Scenario loaded. Controlled agents (current): {len(env.controlled_agents)}, total in scenario: {env.num_controlled_in_scenario}")
|
||||||
|
|
||||||
for step in range(args.horizon):
|
for step in range(args.horizon):
|
||||||
obs, rewards, dones, infos = env.step(None)
|
obs, rewards, dones, infos = env.step(None)
|
||||||
@@ -67,6 +69,7 @@ def _run_replay(args):
|
|||||||
if dones["__all__"]:
|
if dones["__all__"]:
|
||||||
print(f"Scenario {i} finished at step {step}")
|
print(f"Scenario {i} finished at step {step}")
|
||||||
break
|
break
|
||||||
|
env.close()
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
print("Interrupted by user")
|
print("Interrupted by user")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -74,7 +77,6 @@ def _run_replay(args):
|
|||||||
traceback.print_exc()
|
traceback.print_exc()
|
||||||
print(f"Global error: {e}")
|
print(f"Global error: {e}")
|
||||||
finally:
|
finally:
|
||||||
env.close()
|
|
||||||
print("Environment closed.")
|
print("Environment closed.")
|
||||||
|
|
||||||
|
|
||||||
@@ -176,7 +178,10 @@ def _run_policy(args):
|
|||||||
pass
|
pass
|
||||||
continue
|
continue
|
||||||
|
|
||||||
print(f"Scenario loaded. Controlled agents: {len(obs_dict)}")
|
print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {env.num_controlled_in_scenario}")
|
||||||
|
if len(obs_dict) == 0:
|
||||||
|
print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.")
|
||||||
|
continue
|
||||||
step_count = 0
|
step_count = 0
|
||||||
episode_reward = 0.0
|
episode_reward = 0.0
|
||||||
|
|
||||||
|
|||||||
114
train_bc.py
114
train_bc.py
@@ -15,11 +15,12 @@ from torch.utils.tensorboard import SummaryWriter
|
|||||||
from Algorithm.policy import StateIndependentPolicy
|
from Algorithm.policy import StateIndependentPolicy
|
||||||
from Algorithm.bc import train_bc_epoch, eval_bc_epoch
|
from Algorithm.bc import train_bc_epoch, eval_bc_epoch
|
||||||
from Env.bc_env import BCScenarioEnv
|
from Env.bc_env import BCScenarioEnv
|
||||||
from dataset.loader import load_expert_pkl
|
from dataset.loader import load_expert_pkl, get_expert_scenario_ids
|
||||||
|
|
||||||
|
|
||||||
def evaluate_policy(policy, args, device):
|
def evaluate_policy(policy, args, device):
|
||||||
"""在 BCScenarioEnv 中评估策略,跑若干 episode,返回平均 reward。"""
|
"""在 BCScenarioEnv 中评估策略:仅使用专家数据中出现过的 scenario_id,保证 eval 有受控车。
|
||||||
|
输出与 replay 对齐:agents (current)=reset 时受控车数,total in scenario=该场景受控轨迹总数(car_birth_info_list 长度)。"""
|
||||||
waymo_data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
|
waymo_data_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
|
||||||
data_dir = os.path.join(waymo_data_dir, "exp_filtered")
|
data_dir = os.path.join(waymo_data_dir, "exp_filtered")
|
||||||
if not os.path.exists(data_dir):
|
if not os.path.exists(data_dir):
|
||||||
@@ -28,53 +29,74 @@ def evaluate_policy(policy, args, device):
|
|||||||
print(f"[ERROR] Could not find scenario data in {waymo_data_dir}. Evaluation skipped.")
|
print(f"[ERROR] Could not find scenario data in {waymo_data_dir}. Evaluation skipped.")
|
||||||
return 0.0
|
return 0.0
|
||||||
|
|
||||||
env_config = {
|
scenario_ids = get_expert_scenario_ids(args.expert_data_path, max_ids=5)
|
||||||
"data_directory": data_dir,
|
if not scenario_ids:
|
||||||
"is_multi_agent": True,
|
print("[WARN] No scenario_id in expert pkl, falling back to scenarios [0,1,2]. Eval may have 0 controlled agents.")
|
||||||
"num_controlled_agents": 3,
|
scenario_ids = [0, 1, 2]
|
||||||
"use_render": False,
|
|
||||||
"sequential_seed": True,
|
|
||||||
"horizon": 200,
|
|
||||||
}
|
|
||||||
env = BCScenarioEnv(env_config, agent2policy=None)
|
|
||||||
total_rewards = []
|
|
||||||
|
|
||||||
try:
|
total_rewards = []
|
||||||
for i in range(3):
|
horizon = 200
|
||||||
obs_dict = env.reset(seed=i)
|
|
||||||
episode_reward = 0
|
for idx, scenario_id in enumerate(scenario_ids):
|
||||||
dones = {"__all__": False}
|
env_config = {
|
||||||
step_count = 0
|
"data_directory": data_dir,
|
||||||
horizon = 200
|
"is_multi_agent": True,
|
||||||
while not dones["__all__"]:
|
"num_controlled_agents": 100,
|
||||||
step_count += 1
|
"use_render": False,
|
||||||
if step_count >= horizon:
|
"sequential_seed": True,
|
||||||
break
|
"horizon": horizon,
|
||||||
if not obs_dict:
|
"start_scenario_index": scenario_id,
|
||||||
obs_dict, _, dones, _ = env.step({})
|
"num_scenarios": 1,
|
||||||
continue
|
}
|
||||||
agent_ids = list(obs_dict.keys())
|
env = BCScenarioEnv(env_config, agent2policy=None)
|
||||||
obs_list = [obs_dict[aid] for aid in agent_ids]
|
try:
|
||||||
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
obs_dict = env.reset(seed=scenario_id)
|
||||||
with torch.no_grad():
|
except Exception as e:
|
||||||
actions, _ = policy.sample(obs_tensor)
|
print(f" Eval Episode {idx} (scenario {scenario_id}): reset failed: {e}")
|
||||||
actions = actions.cpu().numpy()
|
env.close()
|
||||||
action_dict = {aid: act for aid, act in zip(agent_ids, actions)}
|
continue
|
||||||
obs_dict, rewards, dones, _ = env.step(action_dict)
|
|
||||||
episode_reward += sum(rewards.values())
|
n_controlled = len(env.controlled_agents)
|
||||||
total_rewards.append(episode_reward)
|
n_total_in_scenario = getattr(env, "num_controlled_in_scenario", n_controlled)
|
||||||
print(f" Eval Episode {i}: Total Reward {episode_reward:.2f}")
|
if n_controlled == 0:
|
||||||
avg_reward = float(np.mean(total_rewards))
|
print(
|
||||||
print(f" Average Evaluation Reward: {avg_reward:.2f}")
|
f" Eval Episode {idx} (scenario {scenario_id}): 0 controlled agents (total in scenario: {n_total_in_scenario}), skip."
|
||||||
return avg_reward
|
)
|
||||||
except Exception as e:
|
env.close()
|
||||||
print(f"Evaluation failed: {e}")
|
continue
|
||||||
import traceback
|
|
||||||
traceback.print_exc()
|
episode_reward = 0.0
|
||||||
return 0.0
|
step_count = 0
|
||||||
finally:
|
dones = {"__all__": False}
|
||||||
|
while not dones["__all__"] and step_count < horizon:
|
||||||
|
step_count += 1
|
||||||
|
if not obs_dict:
|
||||||
|
obs_dict, _, dones, _ = env.step({})
|
||||||
|
continue
|
||||||
|
agent_ids = list(obs_dict.keys())
|
||||||
|
obs_list = [obs_dict[aid] for aid in agent_ids]
|
||||||
|
obs_tensor = torch.FloatTensor(np.array(obs_list)).to(device)
|
||||||
|
with torch.no_grad():
|
||||||
|
actions, _ = policy.sample(obs_tensor)
|
||||||
|
actions = actions.cpu().numpy()
|
||||||
|
action_dict = {aid: act for aid, act in zip(agent_ids, actions)}
|
||||||
|
obs_dict, rewards, dones, _ = env.step(action_dict)
|
||||||
|
episode_reward += sum(rewards.values())
|
||||||
|
|
||||||
|
total_rewards.append(episode_reward)
|
||||||
|
print(
|
||||||
|
f" Eval Episode {idx} (scenario {scenario_id}): Total Reward {episode_reward:.2f}, steps {step_count}, "
|
||||||
|
f"agents (current): {n_controlled}, total in scenario: {n_total_in_scenario}"
|
||||||
|
)
|
||||||
env.close()
|
env.close()
|
||||||
|
|
||||||
|
if not total_rewards:
|
||||||
|
print(" No valid eval episodes (all skipped or failed).")
|
||||||
|
return 0.0
|
||||||
|
avg_reward = float(np.mean(total_rewards))
|
||||||
|
print(f" Average Evaluation Reward: {avg_reward:.2f}")
|
||||||
|
return avg_reward
|
||||||
|
|
||||||
|
|
||||||
def main(args):
|
def main(args):
|
||||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|||||||
Reference in New Issue
Block a user