mixed_training
This commit is contained in:
527
Env/expert_replay_env.py
Normal file
527
Env/expert_replay_env.py
Normal file
@@ -0,0 +1,527 @@
|
||||
import logging
|
||||
import numpy as np
|
||||
from collections import defaultdict
|
||||
from metadrive.component.vehicle.vehicle_type import DefaultVehicle
|
||||
from metadrive.type import MetaDriveType
|
||||
from Env.scenario_env import MultiAgentScenarioEnv, PolicyVehicle
|
||||
from Env.inverse_dynamics import InverseDynamics
|
||||
|
||||
class ExpertReplayEnv(MultiAgentScenarioEnv):
|
||||
def __init__(self, config=None):
|
||||
# Allow passing config without agent2policy since we don't use policies for replay
|
||||
if config is None:
|
||||
config = {}
|
||||
# Ensure we don't simulate physics for the controlled agents in the traditional sense
|
||||
# but we still need the engine to run
|
||||
super().__init__(config, agent2policy={})
|
||||
self.inverse_dynamics = InverseDynamics()
|
||||
self.expert_tracks = {}
|
||||
# Replay SDC/ego ("default_agent" in MetaDrive) as well; otherwise it will keep default action=0 and look stuck.
|
||||
self.replay_sdc = self.config.get("replay_sdc", True)
|
||||
self.sdc_track = None
|
||||
self.sdc_vehicle = None
|
||||
self.sdc_agent_id = "default_agent"
|
||||
|
||||
def reset(self, seed=None):
|
||||
self.round = 0
|
||||
if self.logger is None:
|
||||
from metadrive.engine.logger import get_logger, set_log_level
|
||||
self.logger = get_logger()
|
||||
log_level = self.config.get("log_level", logging.INFO)
|
||||
set_log_level(log_level)
|
||||
|
||||
self.lazy_init()
|
||||
self._reset_global_seed(seed)
|
||||
if self.engine is None:
|
||||
raise ValueError("Broken MetaDrive instance.")
|
||||
|
||||
self.background_vehicles = {} # Vehicles that exist but are static/background
|
||||
|
||||
# Helper function to check if a position is on a valid lane
|
||||
def is_on_lane(pos, map_manager, threshold=2.0):
|
||||
# Check if point is close to any lane in the road network
|
||||
# This can be expensive if checked for every point, so we check sample points
|
||||
# or rely on lane index if available.
|
||||
# Waymo tracks don't have lane index, just positions.
|
||||
# We can use map.road_network.get_closest_lane_index(pos)
|
||||
if map_manager is None or map_manager.current_map is None:
|
||||
return True # If no map, assume valid
|
||||
|
||||
try:
|
||||
# Use a larger search radius to catch slightly offset lanes
|
||||
lane, lane_index = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True)
|
||||
if lane is None:
|
||||
return False
|
||||
|
||||
# Check lateral distance
|
||||
long, lat = lane.local_coordinates(pos)
|
||||
width = lane.width
|
||||
# Allow being slightly off-lane (e.g. changing lanes)
|
||||
# But parking lots are usually far from defined lanes in Waymo converted maps
|
||||
if abs(lat) <= (width / 2 + threshold):
|
||||
return True
|
||||
return False
|
||||
except:
|
||||
return False
|
||||
|
||||
# --- MODIFIED SECTION START ---
|
||||
# Capture expert tracks before they are cleaned
|
||||
self.expert_tracks = {}
|
||||
# Capture SDC track for ego replay (MetaDrive default agent)
|
||||
self.sdc_track = None
|
||||
self.sdc_vehicle = None
|
||||
if self.replay_sdc and hasattr(self.engine, "traffic_manager"):
|
||||
sdc_sid = self.engine.traffic_manager.sdc_scenario_id
|
||||
self.sdc_track = self.engine.traffic_manager.current_traffic_data.get(sdc_sid, None)
|
||||
_obj_to_clean_this_frame = []
|
||||
self.car_birth_info_list = []
|
||||
|
||||
# Pre-filter: Check tracks against map AND check for static vehicles
|
||||
|
||||
for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items():
|
||||
if scenario_id == self.engine.traffic_manager.sdc_scenario_id:
|
||||
continue
|
||||
else:
|
||||
if track["type"] == MetaDriveType.VEHICLE:
|
||||
_obj_to_clean_this_frame.append(scenario_id)
|
||||
|
||||
valid = track['state']['valid']
|
||||
if not valid.any():
|
||||
continue
|
||||
|
||||
first_show = np.argmax(valid)
|
||||
last_show = len(valid) - 1 - np.argmax(valid[::-1])
|
||||
mid_show = (first_show + last_show) // 2
|
||||
|
||||
# 1. Lane check (existing logic)
|
||||
points_to_check = [first_show, mid_show, last_show]
|
||||
on_road_count = 0
|
||||
is_valid_track = True
|
||||
start_pos = track['state']['position'][first_show]
|
||||
if not is_on_lane(start_pos, self.engine.map_manager, threshold=5.0): # 5m tolerance
|
||||
mid_pos = track['state']['position'][mid_show]
|
||||
if not is_on_lane(mid_pos, self.engine.map_manager, threshold=5.0):
|
||||
is_valid_track = False
|
||||
|
||||
# 2. Static check
|
||||
# Calculate total displacement and max speed
|
||||
positions = track['state']['position'][valid.astype(bool)]
|
||||
velocities = track['state']['velocity'][valid.astype(bool)]
|
||||
|
||||
total_displacement = 0
|
||||
max_speed = 0
|
||||
if len(positions) > 1:
|
||||
total_displacement = np.linalg.norm(positions[-1] - positions[0])
|
||||
max_speed = np.max(np.linalg.norm(velocities, axis=1))
|
||||
|
||||
is_static = False
|
||||
if total_displacement < 5.0 and max_speed < 1.0: # Relaxed threshold: <5m move and <1m/s
|
||||
is_static = True
|
||||
|
||||
# Decision logic:
|
||||
# - If off-road AND static: Skip completely (don't even spawn as background)
|
||||
# - If off-road but moving: Maybe keep? Or skip? Usually off-road moving is weird, skip.
|
||||
# - If on-road but static: Spawn as BACKGROUND (visible but not controlled agent)
|
||||
# - If on-road and moving: Spawn as CONTROLLED agent
|
||||
|
||||
if not is_valid_track:
|
||||
# Skip off-road vehicles entirely (both static and moving off-road)
|
||||
continue
|
||||
|
||||
if is_static:
|
||||
# Add to background list, but NOT to car_birth_info_list (which is for controlled agents)
|
||||
# We need a way to spawn them. Let's add a separate list.
|
||||
self.background_vehicles[scenario_id] = {
|
||||
'id': track['metadata']['object_id'],
|
||||
'show_time': first_show,
|
||||
'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]),
|
||||
'heading': track['state']['heading'][first_show],
|
||||
'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]),
|
||||
'scenario_id': scenario_id,
|
||||
'length': track['state']['length'][first_show],
|
||||
'width': track['state']['width'][first_show],
|
||||
'valid': valid # Need validity to know when to show/hide
|
||||
}
|
||||
continue # Do not add to controlled list
|
||||
|
||||
# Store the full track for replay (only for controlled agents)
|
||||
self.expert_tracks[scenario_id] = track
|
||||
|
||||
self.car_birth_info_list.append({
|
||||
'id': track['metadata']['object_id'],
|
||||
'show_time': first_show,
|
||||
'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]),
|
||||
'heading': track['state']['heading'][first_show],
|
||||
'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]),
|
||||
'scenario_id': scenario_id, # Keep track of original ID to lookup tracks
|
||||
'length': track['state']['length'][first_show],
|
||||
'width': track['state']['width'][first_show]
|
||||
})
|
||||
|
||||
for scenario_id in _obj_to_clean_this_frame:
|
||||
self.engine.traffic_manager.current_traffic_data.pop(scenario_id)
|
||||
# --- MODIFIED SECTION END ---
|
||||
|
||||
self.engine.reset()
|
||||
self.reset_sensors()
|
||||
self.engine.taskMgr.step()
|
||||
|
||||
self.lanes = self.engine.map_manager.current_map.road_network.graph
|
||||
|
||||
if self.top_down_renderer is not None:
|
||||
self.top_down_renderer.clear()
|
||||
self.engine.top_down_renderer = None
|
||||
|
||||
self.dones = {}
|
||||
self.episode_rewards = defaultdict(float)
|
||||
self.episode_lengths = defaultdict(int)
|
||||
|
||||
self.controlled_agents.clear()
|
||||
self.controlled_agent_ids.clear()
|
||||
|
||||
# We skip calling super().reset() to avoid double reset
|
||||
# But we need to ensure ScenarioEnv-specific setup is done if any.
|
||||
# ScenarioEnv.reset() basically does engine.reset() and some cleanup.
|
||||
# We covered most of it.
|
||||
|
||||
self._spawn_controlled_agents()
|
||||
self._spawn_background_vehicles() # Initial spawn for background
|
||||
|
||||
# Ensure SDC/ego is moved to the correct initial expert state.
|
||||
if self.replay_sdc:
|
||||
self.sdc_vehicle = self.engine.agent_manager.active_agents.get(self.sdc_agent_id, None)
|
||||
if self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||
valid = self.sdc_track["state"]["valid"]
|
||||
t0 = int(np.argmax(valid)) if valid.any() else 0
|
||||
pos0 = self.sdc_track["state"]["position"][t0]
|
||||
heading0 = self.sdc_track["state"]["heading"][t0]
|
||||
vel0 = self.sdc_track["state"]["velocity"][t0]
|
||||
self.sdc_vehicle.set_position(pos0)
|
||||
self.sdc_vehicle.set_heading_theta(heading0)
|
||||
self.sdc_vehicle.set_velocity(vel0)
|
||||
|
||||
return self._get_all_obs()
|
||||
|
||||
def _spawn_background_vehicles(self):
|
||||
# Spawn static/background vehicles
|
||||
# Since they are static, we might just spawn them once if their show_time is 0
|
||||
# But Waymo tracks have valid bits, they might appear/disappear.
|
||||
# For optimization, if they are truly static (never move), we just spawn them when show_time matches.
|
||||
|
||||
# We need to track spawned background vehicles to remove them if they become invalid?
|
||||
# Since we defined them as "static", they probably stay put.
|
||||
# But validity might change (e.g. late spawn).
|
||||
|
||||
# For simplicity in this step, let's just iterate and spawn if time matches
|
||||
for sid, car in self.background_vehicles.items():
|
||||
if car['show_time'] == self.round:
|
||||
# Spawn as a Traffic Vehicle (not PolicyVehicle), or just a static object?
|
||||
# Using DefaultVehicle is fine, but don't add to controlled_agents
|
||||
|
||||
# Check duplication
|
||||
bg_id = f"bg_{car['id']}"
|
||||
# if bg_id in self.engine.obj_to_id: # obj_to_id might not be available in all versions
|
||||
if bg_id in self.engine.agent_manager.active_agents:
|
||||
continue
|
||||
|
||||
vehicle_config = {}
|
||||
if 'length' in car and 'width' in car:
|
||||
vehicle_config = {
|
||||
"length": car['length'],
|
||||
"width": car['width']
|
||||
}
|
||||
|
||||
v = self.engine.spawn_object(
|
||||
DefaultVehicle,
|
||||
name=bg_id,
|
||||
vehicle_config=vehicle_config,
|
||||
position=car['begin'],
|
||||
heading=car['heading']
|
||||
)
|
||||
|
||||
# Set color to grey/dark to indicate background
|
||||
v.set_velocity([0, 0])
|
||||
# Maybe set color? MetaDrive vehicles random color.
|
||||
# v.set_color(...) if supported
|
||||
|
||||
# Register as an active object but NOT controlled agent
|
||||
# The engine manages it.
|
||||
# CRITICAL: We need it in self.engine.agent_manager.active_agents for Observation?
|
||||
# If we want it to be seen by Lidar/Observation, it needs to be an "agent" or "traffic".
|
||||
# DefaultVehicle spawned this way is just an object.
|
||||
# We should add it to traffic manager? Or just leave it as object?
|
||||
# MultiAgentScenarioEnv._get_all_obs iterates self.engine.agent_manager.active_agents
|
||||
|
||||
# If we want it in observation, we must add it to active_agents OR iterate over all objects.
|
||||
# Adding to active_agents is easier for compatibility.
|
||||
self.engine.agent_manager.active_agents[bg_id] = v
|
||||
|
||||
# Store valid mask to remove it later if needed?
|
||||
v.valid_mask = car['valid']
|
||||
v.start_t = car['show_time']
|
||||
|
||||
def _update_background_vehicles(self):
|
||||
# Remove background vehicles if they become invalid
|
||||
# Or spawn new ones
|
||||
self._spawn_background_vehicles()
|
||||
|
||||
# Check validity for existing
|
||||
to_remove = []
|
||||
for aid, v in self.engine.agent_manager.active_agents.items():
|
||||
if aid.startswith("bg_"):
|
||||
# Check validity
|
||||
if hasattr(v, 'valid_mask'):
|
||||
curr_step = self.round
|
||||
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
||||
to_remove.append(aid)
|
||||
|
||||
for aid in to_remove:
|
||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||
# if aid in self.engine.obj_to_id:
|
||||
# self.engine.clear_objects([self.engine.obj_to_id[aid]])
|
||||
# Instead, we should find the object by ID and clear it.
|
||||
# Since we don't track obj directly, we can't easily clear it without obj ref.
|
||||
# Wait, active_agents stores the vehicle object.
|
||||
# So we can just clear that object.
|
||||
pass
|
||||
|
||||
# Re-iterate to clear objects properly
|
||||
for aid in to_remove:
|
||||
# We need to find the vehicle object to clear it.
|
||||
# But we popped it from active_agents.
|
||||
# Wait, we should get it before pop.
|
||||
pass
|
||||
|
||||
def _update_background_vehicles(self):
|
||||
# Remove background vehicles if they become invalid
|
||||
# Or spawn new ones
|
||||
self._spawn_background_vehicles()
|
||||
|
||||
# Check validity for existing
|
||||
to_remove = []
|
||||
objects_to_clear = []
|
||||
|
||||
for aid, v in self.engine.agent_manager.active_agents.items():
|
||||
if aid.startswith("bg_"):
|
||||
# Check validity
|
||||
if hasattr(v, 'valid_mask'):
|
||||
curr_step = self.round
|
||||
if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]:
|
||||
to_remove.append(aid)
|
||||
objects_to_clear.append(v)
|
||||
|
||||
for aid in to_remove:
|
||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||
|
||||
if objects_to_clear:
|
||||
self.engine.clear_objects(objects_to_clear)
|
||||
|
||||
def _spawn_controlled_agents(self):
|
||||
for car in self.car_birth_info_list:
|
||||
if car['show_time'] == self.round:
|
||||
agent_id = f"controlled_{car['id']}"
|
||||
|
||||
# Check if we already have this agent (shouldn't happen with unique IDs but safety check)
|
||||
if agent_id in self.controlled_agents:
|
||||
continue
|
||||
|
||||
# Handling ID flickering / merging
|
||||
# If this ID is new, check if there's an existing agent very close to its start position
|
||||
# that just disappeared? (Not implemented here, complex logic)
|
||||
# But we can check if there's an overlap with existing agents?
|
||||
# For now, just spawn.
|
||||
|
||||
# Read vehicle type/size if available
|
||||
vehicle_config = {}
|
||||
if 'length' in car and 'width' in car:
|
||||
vehicle_config = {
|
||||
"length": car['length'],
|
||||
"width": car['width']
|
||||
}
|
||||
|
||||
vehicle = self.engine.spawn_object(
|
||||
PolicyVehicle,
|
||||
name=agent_id,
|
||||
vehicle_config=vehicle_config,
|
||||
position=car['begin'],
|
||||
heading=car['heading']
|
||||
)
|
||||
vehicle.reset(position=car['begin'], heading=car['heading'])
|
||||
|
||||
# We don't set policy or destination in the same way, or maybe we do for compatibility
|
||||
vehicle.set_destination(car['end'])
|
||||
|
||||
# Store extra info for replay
|
||||
vehicle.expert_track = self.expert_tracks[car['scenario_id']]
|
||||
vehicle.original_id = car['id']
|
||||
|
||||
self.controlled_agents[agent_id] = vehicle
|
||||
self.controlled_agent_ids.append(agent_id)
|
||||
self.engine.agent_manager.active_agents[agent_id] = vehicle
|
||||
|
||||
def step(self, action_dict=None):
|
||||
# We ignore input action_dict for the purpose of controlling agents
|
||||
# Instead, we calculate what the action *should* be
|
||||
|
||||
self.round += 1
|
||||
expert_actions = {}
|
||||
|
||||
# 1. Update state of all controlled agents to the current timestep (self.round)
|
||||
# and compute action from (self.round-1) to (self.round).
|
||||
# Wait, usually step() moves T -> T+1.
|
||||
# Current state is T. We want to move to T+1.
|
||||
# So we need state at T and T+1.
|
||||
|
||||
# Identify agents that are done (valid=0 at T+1 or T+1 >= length)
|
||||
agents_to_remove = []
|
||||
|
||||
# Update SDC/ego first (otherwise it will stay still with default action=0)
|
||||
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||
next_step = self.round
|
||||
curr_step = self.round - 1
|
||||
if next_step < len(self.sdc_track["state"]["position"]) and self.sdc_track["state"]["valid"][next_step]:
|
||||
curr_state = {
|
||||
"position": self.sdc_track["state"]["position"][curr_step],
|
||||
"heading": self.sdc_track["state"]["heading"][curr_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][curr_step],
|
||||
}
|
||||
next_state = {
|
||||
"position": self.sdc_track["state"]["position"][next_step],
|
||||
"heading": self.sdc_track["state"]["heading"][next_step],
|
||||
"velocity": self.sdc_track["state"]["velocity"][next_step],
|
||||
}
|
||||
action, _ = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1)
|
||||
expert_actions[self.sdc_agent_id] = action
|
||||
self.sdc_vehicle.set_position(next_state["position"])
|
||||
self.sdc_vehicle.set_heading_theta(next_state["heading"])
|
||||
self.sdc_vehicle.set_velocity(next_state["velocity"])
|
||||
self.sdc_vehicle.last_expert_action = action
|
||||
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
track = vehicle.expert_track
|
||||
# current_step = self.round - 1 # Since we incremented at start
|
||||
# But vehicle is currently at state corresponding to self.round - 1.
|
||||
# We want to move it to self.round.
|
||||
|
||||
# Check bounds
|
||||
next_step = self.round
|
||||
curr_step = self.round - 1
|
||||
|
||||
if next_step >= len(track['state']['position']):
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
|
||||
valid = track['state']['valid'][next_step]
|
||||
if not valid:
|
||||
agents_to_remove.append(agent_id)
|
||||
continue
|
||||
|
||||
# Get states
|
||||
curr_pos = track['state']['position'][curr_step]
|
||||
next_pos = track['state']['position'][next_step]
|
||||
curr_heading = track['state']['heading'][curr_step]
|
||||
next_heading = track['state']['heading'][next_step]
|
||||
curr_vel = track['state']['velocity'][curr_step]
|
||||
next_vel = track['state']['velocity'][next_step]
|
||||
|
||||
# Prepare state dicts for Inverse Dynamics
|
||||
curr_state = {
|
||||
'position': curr_pos,
|
||||
'heading': curr_heading,
|
||||
'velocity': curr_vel
|
||||
}
|
||||
next_state = {
|
||||
'position': next_pos,
|
||||
'heading': next_heading,
|
||||
'velocity': next_vel
|
||||
}
|
||||
|
||||
# Calculate action
|
||||
action, raw_info = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1) # Waymo is 10Hz?
|
||||
expert_actions[agent_id] = action
|
||||
|
||||
# Force update vehicle state
|
||||
vehicle.set_position(next_pos)
|
||||
vehicle.set_heading_theta(next_heading)
|
||||
vehicle.set_velocity(next_vel)
|
||||
|
||||
# Also record this action in the vehicle for later retrieval if needed
|
||||
vehicle.last_expert_action = action
|
||||
|
||||
# Remove finished agents
|
||||
for agent_id in agents_to_remove:
|
||||
vehicle = self.controlled_agents[agent_id]
|
||||
self.controlled_agents.pop(agent_id)
|
||||
self.controlled_agent_ids.remove(agent_id)
|
||||
self.engine.agent_manager.active_agents.pop(agent_id, None)
|
||||
|
||||
self.engine.clear_objects([vehicle.id])
|
||||
|
||||
# Step physics world to update sensors/collision detection
|
||||
# We don't need full integration, but we need to update the physics world state
|
||||
self.engine.taskMgr.step()
|
||||
|
||||
# Spawn new agents for this turn
|
||||
self._spawn_controlled_agents()
|
||||
self._update_background_vehicles()
|
||||
|
||||
# Get observations
|
||||
obs = self._get_all_obs()
|
||||
|
||||
rewards = {aid: 0.0 for aid in self.controlled_agents}
|
||||
dones = {aid: False for aid in self.controlled_agents}
|
||||
dones["__all__"] = (self.round >= self.config["horizon"]) or (len(self.controlled_agents) == 0 and self.round > 190) # Waymo scenarios are usually ~198 steps (20s @ 10Hz) or 90 steps (9s)
|
||||
|
||||
infos = {aid: {"expert_action": expert_actions.get(aid, np.zeros(2))} for aid in self.controlled_agents}
|
||||
|
||||
return obs, rewards, dones, infos
|
||||
|
||||
def _get_all_obs(self):
|
||||
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||
obs_dict = {}
|
||||
|
||||
for agent_id, vehicle in self.controlled_agents.items():
|
||||
# 1. Ego State
|
||||
ego_state = [
|
||||
vehicle.position[0], vehicle.position[1],
|
||||
vehicle.velocity[0], vehicle.velocity[1],
|
||||
vehicle.heading_theta
|
||||
]
|
||||
|
||||
# 2. Neighbors
|
||||
neighbors = []
|
||||
# Iterate through all vehicles in the engine
|
||||
candidates = []
|
||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||
if other_id == agent_id:
|
||||
continue
|
||||
|
||||
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||
if dist < 30.0:
|
||||
candidates.append((dist, other_vehicle))
|
||||
|
||||
# Sort by distance
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
|
||||
# Take top 10
|
||||
top_10 = candidates[:10]
|
||||
|
||||
neighbor_feats = []
|
||||
for _, neighbor in top_10:
|
||||
neighbor_feats.extend([
|
||||
neighbor.position[0] - vehicle.position[0], # Relative pos
|
||||
neighbor.position[1] - vehicle.position[1],
|
||||
neighbor.velocity[0], # Absolute vel? or Relative? Usually relative in MultiAgent
|
||||
neighbor.velocity[1]
|
||||
])
|
||||
|
||||
# Pad if < 10
|
||||
missing = 10 - len(top_10)
|
||||
if missing > 0:
|
||||
neighbor_feats.extend([0.0] * (4 * missing))
|
||||
|
||||
# Flatten
|
||||
obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||
obs_dict[agent_id] = obs
|
||||
|
||||
return obs_dict
|
||||
Reference in New Issue
Block a user