HBBC部署到代码中
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
194
Env/bc_ego_replay_env.py
Normal file
194
Env/bc_ego_replay_env.py
Normal file
@@ -0,0 +1,194 @@
|
|||||||
|
"""
|
||||||
|
Single-agent BC evaluation environment: only ego (SDC) is controlled by the policy;
|
||||||
|
other vehicles are replayed from expert trajectories (same as data collection).
|
||||||
|
"""
|
||||||
|
import numpy as np
|
||||||
|
from Env.expert_replay_env import ExpertReplayEnv
|
||||||
|
from Env.hbbc_background_policy import HBBCBackgroundController
|
||||||
|
|
||||||
|
|
||||||
|
class BCEgoReplayEnv(ExpertReplayEnv):
|
||||||
|
"""
|
||||||
|
For single-agent BC evaluation: controlled_agents exposes only SDC (default_agent).
|
||||||
|
Other vehicles are still spawned and replayed by expert; internally we keep them
|
||||||
|
in _replay_agents so step() can update them.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def reset(self, seed=None):
|
||||||
|
obs = super().reset(seed=seed)
|
||||||
|
self.enable_hbbc_background = bool(self.config.get("enable_hbbc_background", False))
|
||||||
|
self.hbbc_controller = None
|
||||||
|
self._hbbc_runtime_logged = False
|
||||||
|
if self.enable_hbbc_background:
|
||||||
|
self.hbbc_controller = HBBCBackgroundController(
|
||||||
|
model_path=self.config.get("hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||||
|
device=self.config.get("hbbc_inference_device", "cpu"),
|
||||||
|
latent_mode=self.config.get("hbbc_latent_mode", "per_vehicle_fixed"),
|
||||||
|
latent_json_path=self.config.get("hbbc_latent_json_path"),
|
||||||
|
seed=int(self.config.get("seed", 0)),
|
||||||
|
dt=float(self.config.get("hbbc_dt", 0.1)),
|
||||||
|
)
|
||||||
|
self.hbbc_controller.reset_episode()
|
||||||
|
# Expose only SDC as the controlled agent for the evaluator
|
||||||
|
self._replay_agents = dict(self.controlled_agents)
|
||||||
|
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||||
|
self.controlled_agents = {self.sdc_agent_id: self.sdc_vehicle}
|
||||||
|
self.controlled_agent_ids = [self.sdc_agent_id]
|
||||||
|
else:
|
||||||
|
self.controlled_agents = {}
|
||||||
|
self.controlled_agent_ids = []
|
||||||
|
return self._get_all_obs()
|
||||||
|
|
||||||
|
def _get_all_obs(self):
|
||||||
|
"""Return only ego (SDC) observation so evaluator has a single agent."""
|
||||||
|
if not self.controlled_agents or self.sdc_vehicle is None:
|
||||||
|
return {}
|
||||||
|
obs = self._obs_for_vehicle(self.sdc_vehicle, exclude_agent_id=self.sdc_agent_id)
|
||||||
|
return {self.sdc_agent_id: obs}
|
||||||
|
|
||||||
|
def step(self, action_dict=None):
|
||||||
|
self.round += 1
|
||||||
|
expert_actions = {}
|
||||||
|
agents_to_remove = []
|
||||||
|
|
||||||
|
# SDC: use policy action if provided, else expert replay
|
||||||
|
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_track is not None:
|
||||||
|
policy_action = None
|
||||||
|
if action_dict and self.sdc_agent_id in action_dict:
|
||||||
|
policy_action = np.asarray(action_dict[self.sdc_agent_id], dtype=np.float64)
|
||||||
|
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],
|
||||||
|
}
|
||||||
|
if policy_action is not None:
|
||||||
|
next_state = self.inverse_dynamics.apply_action(curr_state, policy_action, dt=0.1)
|
||||||
|
expert_actions[self.sdc_agent_id] = policy_action
|
||||||
|
else:
|
||||||
|
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 = expert_actions[self.sdc_agent_id]
|
||||||
|
|
||||||
|
# Replay other vehicles: restore full controlled_agents for internal logic
|
||||||
|
self.controlled_agents = dict(self._replay_agents)
|
||||||
|
self.controlled_agent_ids = list(self.controlled_agents.keys())
|
||||||
|
hbbc_batch = []
|
||||||
|
hbbc_curr_states = {}
|
||||||
|
for agent_id, vehicle in self.controlled_agents.items():
|
||||||
|
track = vehicle.expert_track
|
||||||
|
next_step = self.round
|
||||||
|
if next_step >= len(track["state"]["position"]):
|
||||||
|
agents_to_remove.append(agent_id)
|
||||||
|
continue
|
||||||
|
if not track["state"]["valid"][next_step]:
|
||||||
|
agents_to_remove.append(agent_id)
|
||||||
|
continue
|
||||||
|
if self.enable_hbbc_background and self.hbbc_controller is not None:
|
||||||
|
# HBBC autonomous rollout: use vehicle's own previous-step state
|
||||||
|
curr_state = {
|
||||||
|
"position": np.asarray(vehicle.position, dtype=np.float64),
|
||||||
|
"heading": float(vehicle.heading_theta),
|
||||||
|
"velocity": np.asarray(vehicle.velocity, dtype=np.float64),
|
||||||
|
}
|
||||||
|
object_id = str(getattr(vehicle, "original_id", agent_id))
|
||||||
|
hbbc_batch.append((agent_id, vehicle, object_id, agent_id))
|
||||||
|
hbbc_curr_states[agent_id] = curr_state
|
||||||
|
else:
|
||||||
|
curr_step = self.round - 1
|
||||||
|
curr_state = {
|
||||||
|
"position": track["state"]["position"][curr_step],
|
||||||
|
"heading": track["state"]["heading"][curr_step],
|
||||||
|
"velocity": track["state"]["velocity"][curr_step],
|
||||||
|
}
|
||||||
|
next_state = {
|
||||||
|
"position": track["state"]["position"][next_step],
|
||||||
|
"heading": track["state"]["heading"][next_step],
|
||||||
|
"velocity": track["state"]["velocity"][next_step],
|
||||||
|
}
|
||||||
|
action, _ = self.inverse_dynamics.compute_action(curr_state, next_state, dt=0.1)
|
||||||
|
expert_actions[agent_id] = action
|
||||||
|
vehicle.set_position(next_state["position"])
|
||||||
|
vehicle.set_heading_theta(next_state["heading"])
|
||||||
|
vehicle.set_velocity(next_state["velocity"])
|
||||||
|
vehicle.last_expert_action = action
|
||||||
|
|
||||||
|
if hbbc_batch and self.hbbc_controller is not None:
|
||||||
|
hbbc_actions = self.hbbc_controller.infer_actions(hbbc_batch)
|
||||||
|
if not self._hbbc_runtime_logged:
|
||||||
|
print(f"[HBBC] background policy active, current dynamic agents: {len(hbbc_batch)}")
|
||||||
|
self._hbbc_runtime_logged = True
|
||||||
|
for agent_id, _, _, _ in hbbc_batch:
|
||||||
|
curr_state = hbbc_curr_states[agent_id]
|
||||||
|
action = hbbc_actions[agent_id]
|
||||||
|
next_state = self.inverse_dynamics.apply_action(curr_state, action, dt=0.1)
|
||||||
|
expert_actions[agent_id] = action
|
||||||
|
vehicle = self.controlled_agents[agent_id]
|
||||||
|
vehicle.set_position(next_state["position"])
|
||||||
|
vehicle.set_heading_theta(next_state["heading"])
|
||||||
|
vehicle.set_velocity(next_state["velocity"])
|
||||||
|
try:
|
||||||
|
vehicle.last_current_action.append(action)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
vehicle.last_expert_action = action
|
||||||
|
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])
|
||||||
|
if self.hbbc_controller is not None:
|
||||||
|
self.hbbc_controller.remove_vehicle(agent_id)
|
||||||
|
self.engine.taskMgr.step()
|
||||||
|
self._spawn_controlled_agents()
|
||||||
|
self._update_background_vehicles()
|
||||||
|
self._replay_agents = dict(self.controlled_agents)
|
||||||
|
# Expose only SDC again
|
||||||
|
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||||
|
self.controlled_agents = {self.sdc_agent_id: self.sdc_vehicle}
|
||||||
|
self.controlled_agent_ids = [self.sdc_agent_id]
|
||||||
|
else:
|
||||||
|
self.controlled_agents = {}
|
||||||
|
self.controlled_agent_ids = []
|
||||||
|
|
||||||
|
obs = self._get_all_obs()
|
||||||
|
rewards = {}
|
||||||
|
infos = {aid: {"expert_action": expert_actions.get(aid, np.zeros(2))} for aid in self.controlled_agents}
|
||||||
|
if self.sdc_agent_id in self.controlled_agents and self.sdc_vehicle is not None:
|
||||||
|
speed_coef = float(self.config.get("reward_speed_coef", 0.05))
|
||||||
|
collision_distance = float(self.config.get("collision_distance", 6.0))
|
||||||
|
collision_penalty = float(self.config.get("collision_penalty", 100.0))
|
||||||
|
speed = float(np.linalg.norm(self.sdc_vehicle.velocity))
|
||||||
|
r_speed = speed_coef * speed
|
||||||
|
min_dist = float("inf")
|
||||||
|
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||||
|
if other_id == self.sdc_agent_id:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
d = float(np.linalg.norm(self.sdc_vehicle.position - other_vehicle.position))
|
||||||
|
min_dist = min(min_dist, d)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
near_collision = min_dist < collision_distance
|
||||||
|
r_collision = -collision_penalty if near_collision else 0.0
|
||||||
|
rewards[self.sdc_agent_id] = r_speed + r_collision
|
||||||
|
infos[self.sdc_agent_id].update(
|
||||||
|
near_collision=near_collision,
|
||||||
|
min_dist=min_dist if np.isfinite(min_dist) else None,
|
||||||
|
r_speed=r_speed,
|
||||||
|
r_collision=r_collision,
|
||||||
|
)
|
||||||
|
dones = {aid: False for aid in self.controlled_agents}
|
||||||
|
dones["__all__"] = self.round >= self.config["horizon"] or (len(self._replay_agents) == 0 and self.round > 190)
|
||||||
|
return obs, rewards, dones, infos
|
||||||
@@ -1,4 +1,5 @@
|
|||||||
from Env.scenario_env import MultiAgentScenarioEnv
|
from Env.scenario_env import MultiAgentScenarioEnv
|
||||||
|
from Env.hbbc_background_policy import HBBCBackgroundController
|
||||||
from Env.utils import filter_traffic_tracks_to_birth_lists
|
from Env.utils import filter_traffic_tracks_to_birth_lists
|
||||||
from metadrive.component.vehicle.vehicle_type import DefaultVehicle
|
from metadrive.component.vehicle.vehicle_type import DefaultVehicle
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -15,19 +16,71 @@ class BCScenarioEnv(MultiAgentScenarioEnv):
|
|||||||
expert data is generated with ExpertReplayEnv which includes bg_* in active_agents, so the policy
|
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.
|
was trained on obs that can include those neighbors. Demo should use the same scene for consistency.
|
||||||
"""
|
"""
|
||||||
|
def _init_hbbc_background(self):
|
||||||
|
self.enable_hbbc_background = bool(self.config.get("enable_hbbc_background", False))
|
||||||
|
self.hbbc_dynamic_agents = {}
|
||||||
|
self._spawned_dynamic_bg_ids = set()
|
||||||
|
self.hbbc_controller = None
|
||||||
|
if not self.enable_hbbc_background:
|
||||||
|
return
|
||||||
|
self.hbbc_controller = HBBCBackgroundController(
|
||||||
|
model_path=self.config.get("hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||||
|
device=self.config.get("hbbc_inference_device", "cpu"),
|
||||||
|
latent_mode=self.config.get("hbbc_latent_mode", "per_vehicle_fixed"),
|
||||||
|
latent_json_path=self.config.get("hbbc_latent_json_path"),
|
||||||
|
seed=int(self.config.get("seed", 0)),
|
||||||
|
dt=float(self.config.get("hbbc_dt", 0.1)),
|
||||||
|
)
|
||||||
|
self.hbbc_controller.reset_episode()
|
||||||
|
|
||||||
|
def _move_excess_controlled_to_hbbc_background(self):
|
||||||
|
if not self.enable_hbbc_background:
|
||||||
|
return
|
||||||
|
keep_n = int(self.config.get("num_controlled_agents", 0))
|
||||||
|
keep_n = max(0, keep_n)
|
||||||
|
ordered_ids = list(self.controlled_agents.keys())
|
||||||
|
keep_ids = set(ordered_ids[:keep_n])
|
||||||
|
move_ids = [aid for aid in ordered_ids if aid not in keep_ids]
|
||||||
|
for aid in move_ids:
|
||||||
|
self.hbbc_dynamic_agents[aid] = self.controlled_agents[aid]
|
||||||
|
self.controlled_agents.pop(aid, None)
|
||||||
|
if aid in self.controlled_agent_ids:
|
||||||
|
self.controlled_agent_ids.remove(aid)
|
||||||
|
self._spawned_dynamic_bg_ids.update(move_ids)
|
||||||
|
|
||||||
|
def _apply_hbbc_before_step(self):
|
||||||
|
if not self.enable_hbbc_background or not self.hbbc_dynamic_agents:
|
||||||
|
return
|
||||||
|
batch = []
|
||||||
|
for aid, vehicle in self.hbbc_dynamic_agents.items():
|
||||||
|
object_id = getattr(vehicle, "original_id", None) or aid.replace("controlled_", "", 1)
|
||||||
|
batch.append((aid, vehicle, str(object_id) if object_id is not None else None, aid))
|
||||||
|
actions = self.hbbc_controller.infer_actions(batch)
|
||||||
|
for aid, vehicle in self.hbbc_dynamic_agents.items():
|
||||||
|
action = actions.get(aid, np.zeros(2, dtype=np.float32))
|
||||||
|
vehicle.before_step(action)
|
||||||
|
|
||||||
|
def _apply_hbbc_after_step(self):
|
||||||
|
if not self.enable_hbbc_background:
|
||||||
|
return
|
||||||
|
for vehicle in self.hbbc_dynamic_agents.values():
|
||||||
|
vehicle.after_step()
|
||||||
|
|
||||||
def reset(self, seed=None):
|
def reset(self, seed=None):
|
||||||
|
self._init_hbbc_background()
|
||||||
# Clear background vehicles from previous episode so engine.reset() passes _object_clean_check
|
# Clear background vehicles from previous episode so engine.reset() passes _object_clean_check
|
||||||
if getattr(self, "engine", None) is not None:
|
if getattr(self, "engine", None) is not None:
|
||||||
ids_bg = [
|
ids_bg = [
|
||||||
oid for oid, obj in self.engine.get_objects().items()
|
oid for oid, obj in self.engine.get_objects().items()
|
||||||
if (getattr(obj, "name", None) or getattr(obj, "id", None) or "").startswith("bg_")
|
if (getattr(obj, "name", None) or getattr(obj, "id", None) or "").startswith(("bg_", "controlled_"))
|
||||||
]
|
]
|
||||||
if ids_bg:
|
if ids_bg:
|
||||||
self.engine.clear_objects(ids_bg, force_destroy=True)
|
self.engine.clear_objects(ids_bg, force_destroy=True)
|
||||||
for aid in list(self.engine.agent_manager.active_agents.keys()):
|
for aid in list(self.engine.agent_manager.active_agents.keys()):
|
||||||
if aid.startswith("bg_"):
|
if aid.startswith("bg_") or aid.startswith("controlled_"):
|
||||||
self.engine.agent_manager.active_agents.pop(aid, None)
|
self.engine.agent_manager.active_agents.pop(aid, None)
|
||||||
obs = super().reset(seed=seed)
|
obs = super().reset(seed=seed)
|
||||||
|
self._move_excess_controlled_to_hbbc_background()
|
||||||
self._spawn_background_vehicles()
|
self._spawn_background_vehicles()
|
||||||
return self._get_all_obs()
|
return self._get_all_obs()
|
||||||
|
|
||||||
@@ -76,16 +129,21 @@ class BCScenarioEnv(MultiAgentScenarioEnv):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def step(self, action_dict):
|
def step(self, action_dict):
|
||||||
|
if action_dict is None:
|
||||||
|
action_dict = {}
|
||||||
self.round += 1
|
self.round += 1
|
||||||
for agent_id, action in action_dict.items():
|
for agent_id, action in action_dict.items():
|
||||||
if agent_id in self.controlled_agents:
|
if agent_id in self.controlled_agents:
|
||||||
self.controlled_agents[agent_id].before_step(action)
|
self.controlled_agents[agent_id].before_step(action)
|
||||||
|
self._apply_hbbc_before_step()
|
||||||
self.engine.step()
|
self.engine.step()
|
||||||
self.engine.after_step()
|
self.engine.after_step()
|
||||||
for agent_id in action_dict:
|
for agent_id in action_dict:
|
||||||
if agent_id in self.controlled_agents:
|
if agent_id in self.controlled_agents:
|
||||||
self.controlled_agents[agent_id].after_step()
|
self.controlled_agents[agent_id].after_step()
|
||||||
self._spawn_controlled_agents()
|
self._spawn_controlled_agents()
|
||||||
|
self._move_excess_controlled_to_hbbc_background()
|
||||||
|
self._apply_hbbc_after_step()
|
||||||
self._update_background_vehicles()
|
self._update_background_vehicles()
|
||||||
obs = self._get_all_obs()
|
obs = self._get_all_obs()
|
||||||
|
|
||||||
|
|||||||
@@ -294,60 +294,52 @@ class ExpertReplayEnv(MultiAgentScenarioEnv):
|
|||||||
# Get observations
|
# Get observations
|
||||||
obs = self._get_all_obs()
|
obs = self._get_all_obs()
|
||||||
|
|
||||||
rewards = {aid: 0.0 for aid in self.controlled_agents}
|
# Build rewards/dones/infos: include controlled_agents and optionally SDC for data collection
|
||||||
dones = {aid: False for aid in self.controlled_agents}
|
all_agent_ids = list(self.controlled_agents.keys())
|
||||||
|
if self.replay_sdc and self.sdc_vehicle is not None and self.sdc_agent_id not in all_agent_ids:
|
||||||
|
all_agent_ids = all_agent_ids + [self.sdc_agent_id]
|
||||||
|
rewards = {aid: 0.0 for aid in all_agent_ids}
|
||||||
|
dones = {aid: False for aid in all_agent_ids}
|
||||||
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)
|
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 all_agent_ids}
|
||||||
infos = {aid: {"expert_action": expert_actions.get(aid, np.zeros(2))} for aid in self.controlled_agents}
|
|
||||||
|
|
||||||
return obs, rewards, dones, infos
|
return obs, rewards, dones, infos
|
||||||
|
|
||||||
def _get_all_obs(self):
|
def _obs_for_vehicle(self, vehicle, exclude_agent_id=None):
|
||||||
# Implement custom observation: 30m range, 10 nearest vehicles
|
"""Compute 45-dim obs (ego 5 + 10 neighbors x 4) for a vehicle. exclude_agent_id: do not count as neighbor."""
|
||||||
obs_dict = {}
|
|
||||||
|
|
||||||
for agent_id, vehicle in self.controlled_agents.items():
|
|
||||||
# 1. Ego State
|
|
||||||
ego_state = [
|
ego_state = [
|
||||||
vehicle.position[0], vehicle.position[1],
|
vehicle.position[0], vehicle.position[1],
|
||||||
vehicle.velocity[0], vehicle.velocity[1],
|
vehicle.velocity[0], vehicle.velocity[1],
|
||||||
vehicle.heading_theta
|
vehicle.heading_theta
|
||||||
]
|
]
|
||||||
|
|
||||||
# 2. Neighbors
|
|
||||||
neighbors = []
|
|
||||||
# Iterate through all vehicles in the engine
|
|
||||||
candidates = []
|
candidates = []
|
||||||
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
for other_id, other_vehicle in self.engine.agent_manager.active_agents.items():
|
||||||
if other_id == agent_id:
|
if other_id == exclude_agent_id:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
dist = np.linalg.norm(vehicle.position - other_vehicle.position)
|
||||||
if dist < 30.0:
|
if dist < 30.0:
|
||||||
candidates.append((dist, other_vehicle))
|
candidates.append((dist, other_vehicle))
|
||||||
|
|
||||||
# Sort by distance
|
|
||||||
candidates.sort(key=lambda x: x[0])
|
candidates.sort(key=lambda x: x[0])
|
||||||
|
|
||||||
# Take top 10
|
|
||||||
top_10 = candidates[:10]
|
top_10 = candidates[:10]
|
||||||
|
|
||||||
neighbor_feats = []
|
neighbor_feats = []
|
||||||
for _, neighbor in top_10:
|
for _, neighbor in top_10:
|
||||||
neighbor_feats.extend([
|
neighbor_feats.extend([
|
||||||
neighbor.position[0] - vehicle.position[0], # Relative pos
|
neighbor.position[0] - vehicle.position[0],
|
||||||
neighbor.position[1] - vehicle.position[1],
|
neighbor.position[1] - vehicle.position[1],
|
||||||
neighbor.velocity[0], # Absolute vel? or Relative? Usually relative in MultiAgent
|
neighbor.velocity[0],
|
||||||
neighbor.velocity[1]
|
neighbor.velocity[1]
|
||||||
])
|
])
|
||||||
|
|
||||||
# Pad if < 10
|
|
||||||
missing = 10 - len(top_10)
|
missing = 10 - len(top_10)
|
||||||
if missing > 0:
|
if missing > 0:
|
||||||
neighbor_feats.extend([0.0] * (4 * missing))
|
neighbor_feats.extend([0.0] * (4 * missing))
|
||||||
|
return np.array(ego_state + neighbor_feats, dtype=np.float32)
|
||||||
|
|
||||||
# Flatten
|
def _get_all_obs(self):
|
||||||
obs = np.array(ego_state + neighbor_feats, dtype=np.float32)
|
# Implement custom observation: 30m range, 10 nearest vehicles
|
||||||
obs_dict[agent_id] = obs
|
obs_dict = {}
|
||||||
|
for agent_id, vehicle in self.controlled_agents.items():
|
||||||
|
obs_dict[agent_id] = self._obs_for_vehicle(vehicle, exclude_agent_id=agent_id)
|
||||||
|
# Include SDC/ego obs for expert data collection (e.g. single-agent)
|
||||||
|
if self.replay_sdc and self.sdc_vehicle is not None:
|
||||||
|
obs_dict[self.sdc_agent_id] = self._obs_for_vehicle(self.sdc_vehicle, exclude_agent_id=self.sdc_agent_id)
|
||||||
return obs_dict
|
return obs_dict
|
||||||
|
|||||||
69
Env/hbbc_actor_critic.py
Normal file
69
Env/hbbc_actor_critic.py
Normal file
@@ -0,0 +1,69 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
|
def _get_activation(name: str):
|
||||||
|
name = (name or "elu").lower()
|
||||||
|
mapping = {
|
||||||
|
"elu": nn.ELU,
|
||||||
|
"relu": nn.ReLU,
|
||||||
|
"tanh": nn.Tanh,
|
||||||
|
"leakyrelu": nn.LeakyReLU,
|
||||||
|
}
|
||||||
|
if name not in mapping:
|
||||||
|
raise ValueError(f"Unsupported activation: {name}")
|
||||||
|
return mapping[name]()
|
||||||
|
|
||||||
|
|
||||||
|
class ActorCritic(nn.Module):
|
||||||
|
"""Minimal HBBC ActorCritic for inference-only deployment."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_actor_obs=18,
|
||||||
|
num_critic_obs=18,
|
||||||
|
num_actions=2,
|
||||||
|
latent_c_dim=4,
|
||||||
|
latent_eps_dim=6,
|
||||||
|
use_style_latent=True,
|
||||||
|
actor_hidden_dims=None,
|
||||||
|
activation="elu",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
_ = num_critic_obs # kept for checkpoint compatibility
|
||||||
|
if actor_hidden_dims is None:
|
||||||
|
actor_hidden_dims = [512, 256, 128]
|
||||||
|
|
||||||
|
act_fn = _get_activation(activation)
|
||||||
|
self.latent_c_dim = int(latent_c_dim)
|
||||||
|
self.latent_eps_dim = int(latent_eps_dim)
|
||||||
|
self.use_style_latent = bool(use_style_latent)
|
||||||
|
|
||||||
|
layers = [nn.Linear(num_actor_obs, actor_hidden_dims[0]), act_fn]
|
||||||
|
for i in range(len(actor_hidden_dims) - 1):
|
||||||
|
layers.append(nn.Linear(actor_hidden_dims[i], actor_hidden_dims[i + 1]))
|
||||||
|
layers.append(_get_activation(activation))
|
||||||
|
self.actor_trunk = nn.Sequential(*layers)
|
||||||
|
self.actor_head = nn.Linear(actor_hidden_dims[-1], num_actions)
|
||||||
|
|
||||||
|
if self.use_style_latent:
|
||||||
|
self.style_trunk = nn.Sequential(
|
||||||
|
nn.Linear(self.latent_eps_dim, 512),
|
||||||
|
_get_activation(activation),
|
||||||
|
nn.Linear(512, 256),
|
||||||
|
_get_activation(activation),
|
||||||
|
nn.Linear(256, 128),
|
||||||
|
_get_activation(activation),
|
||||||
|
)
|
||||||
|
self.style_head = nn.Linear(128, self.latent_eps_dim)
|
||||||
|
self.style_activation = torch.tanh
|
||||||
|
|
||||||
|
def act_inference(self, observations: torch.Tensor) -> torch.Tensor:
|
||||||
|
if self.use_style_latent:
|
||||||
|
obs = observations[..., :-(self.latent_c_dim + self.latent_eps_dim)]
|
||||||
|
eps = observations[..., -self.latent_c_dim - self.latent_eps_dim:-self.latent_c_dim]
|
||||||
|
c = observations[..., -self.latent_c_dim:]
|
||||||
|
eps = self.style_activation(self.style_head(self.style_trunk(eps)))
|
||||||
|
observations = torch.cat([obs, eps, c], dim=-1)
|
||||||
|
embedding = self.actor_trunk(observations)
|
||||||
|
return self.actor_head(embedding)
|
||||||
274
Env/hbbc_background_policy.py
Normal file
274
Env/hbbc_background_policy.py
Normal file
@@ -0,0 +1,274 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from Env.hbbc_actor_critic import ActorCritic
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_to_pi(angle: float) -> float:
|
||||||
|
return (angle + np.pi) % (2 * np.pi) - np.pi
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_eps(eps: np.ndarray) -> np.ndarray:
|
||||||
|
eps = np.asarray(eps, dtype=np.float32).reshape(-1)
|
||||||
|
if eps.shape[0] != 6:
|
||||||
|
raise ValueError(f"latent_eps must be 6-dim, got {eps.shape[0]}")
|
||||||
|
norm = float(np.linalg.norm(eps))
|
||||||
|
if norm < 1e-8:
|
||||||
|
eps = np.array([1.0, 0.0, 0.0, 0.0, 0.0, 0.0], dtype=np.float32)
|
||||||
|
else:
|
||||||
|
eps = eps / norm
|
||||||
|
return np.clip(eps, -1.0, 1.0)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_c(latent_c: np.ndarray) -> np.ndarray:
|
||||||
|
c = np.asarray(latent_c, dtype=np.float32).reshape(-1)
|
||||||
|
if c.shape[0] != 4:
|
||||||
|
raise ValueError(f"latent_c must be 4-dim, got {c.shape[0]}")
|
||||||
|
idx = int(np.argmax(c))
|
||||||
|
one_hot = np.zeros(4, dtype=np.float32)
|
||||||
|
one_hot[idx] = 1.0
|
||||||
|
return one_hot
|
||||||
|
|
||||||
|
|
||||||
|
def _sample_latent(rng: np.random.RandomState) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
eps = _normalize_eps(rng.randn(6).astype(np.float32))
|
||||||
|
mode = int(rng.randint(0, 4))
|
||||||
|
c = np.zeros(4, dtype=np.float32)
|
||||||
|
c[mode] = 1.0
|
||||||
|
return eps, c
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class VehicleStateCache:
|
||||||
|
last_heading_theta: Optional[float] = None
|
||||||
|
last_action: Tuple[float, float] = (0.0, 0.0)
|
||||||
|
last_speed_km_h: Optional[float] = None
|
||||||
|
|
||||||
|
|
||||||
|
class HBBCModelWrapper:
|
||||||
|
_cache: Dict[Tuple[str, str], "HBBCModelWrapper"] = {}
|
||||||
|
|
||||||
|
def __init__(self, model_path: str, device: str = "cpu"):
|
||||||
|
self.model_path = os.path.abspath(model_path)
|
||||||
|
self.device = torch.device(device)
|
||||||
|
self.model = self._load_model()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get(cls, model_path: str, device: str = "cpu") -> "HBBCModelWrapper":
|
||||||
|
key = (os.path.abspath(model_path), str(torch.device(device)))
|
||||||
|
if key not in cls._cache:
|
||||||
|
cls._cache[key] = HBBCModelWrapper(model_path=key[0], device=key[1])
|
||||||
|
return cls._cache[key]
|
||||||
|
|
||||||
|
def _load_model(self) -> ActorCritic:
|
||||||
|
model = ActorCritic(
|
||||||
|
num_actor_obs=18,
|
||||||
|
num_critic_obs=18,
|
||||||
|
num_actions=2,
|
||||||
|
latent_c_dim=4,
|
||||||
|
latent_eps_dim=6,
|
||||||
|
use_style_latent=True,
|
||||||
|
).to(self.device)
|
||||||
|
try:
|
||||||
|
ckpt = torch.load(self.model_path, map_location=self.device, weights_only=True)
|
||||||
|
except Exception:
|
||||||
|
ckpt = torch.load(self.model_path, map_location=self.device, weights_only=False)
|
||||||
|
state_dict = ckpt["actor_critic"] if isinstance(ckpt, dict) and "actor_critic" in ckpt else ckpt
|
||||||
|
missing, unexpected = model.load_state_dict(state_dict, strict=False)
|
||||||
|
if missing:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"HBBC checkpoint missing required keys for {self.model_path}: {missing}"
|
||||||
|
)
|
||||||
|
if unexpected:
|
||||||
|
print(f"[HBBC] ignore extra checkpoint keys: {unexpected[:8]}{'...' if len(unexpected) > 8 else ''}")
|
||||||
|
model.eval()
|
||||||
|
return model
|
||||||
|
|
||||||
|
def act_batch(self, obs_batch: np.ndarray) -> np.ndarray:
|
||||||
|
obs_batch = np.asarray(obs_batch, dtype=np.float32)
|
||||||
|
with torch.no_grad():
|
||||||
|
obs_t = torch.from_numpy(obs_batch).to(self.device)
|
||||||
|
actions = self.model.act_inference(obs_t).cpu().numpy()
|
||||||
|
return np.clip(actions, -1.0, 1.0)
|
||||||
|
|
||||||
|
|
||||||
|
class HBBCLatentManager:
|
||||||
|
def __init__(self, mode: str = "per_vehicle_fixed", seed: int = 0, latent_json_path: Optional[str] = None):
|
||||||
|
self.mode = mode
|
||||||
|
self.rng = np.random.RandomState(seed)
|
||||||
|
self.latent_json_path = latent_json_path
|
||||||
|
self.manual_object_latent: Dict[str, Dict[str, np.ndarray]] = {}
|
||||||
|
self.manual_agent_latent: Dict[str, Dict[str, np.ndarray]] = {}
|
||||||
|
self.manual_global_latent: Optional[Tuple[np.ndarray, np.ndarray]] = None
|
||||||
|
self.vehicle_latent: Dict[str, Tuple[np.ndarray, np.ndarray]] = {}
|
||||||
|
self._episode_latent: Optional[Tuple[np.ndarray, np.ndarray]] = None
|
||||||
|
self._load_manual_latent_json()
|
||||||
|
|
||||||
|
def reset_episode(self):
|
||||||
|
self.vehicle_latent.clear()
|
||||||
|
self._episode_latent = None
|
||||||
|
if self.mode == "per_episode_reset":
|
||||||
|
self._episode_latent = _sample_latent(self.rng)
|
||||||
|
|
||||||
|
def _load_manual_latent_json(self):
|
||||||
|
if not self.latent_json_path:
|
||||||
|
return
|
||||||
|
path = os.path.abspath(self.latent_json_path)
|
||||||
|
if not os.path.exists(path):
|
||||||
|
print(f"[HBBC] latent json not found: {path}, fallback to random sampling.")
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[HBBC] failed to load latent json ({path}): {e}. fallback to random sampling.")
|
||||||
|
return
|
||||||
|
|
||||||
|
object_section = data.get("object_id", {})
|
||||||
|
agent_section = data.get("agent_id", {})
|
||||||
|
global_section = data.get("global")
|
||||||
|
|
||||||
|
if global_section is not None:
|
||||||
|
parsed = self._parse_one_latent(global_section, "global")
|
||||||
|
if parsed is not None:
|
||||||
|
self.manual_global_latent = (parsed["latent_eps"], parsed["latent_c"])
|
||||||
|
|
||||||
|
for key, value in object_section.items():
|
||||||
|
parsed = self._parse_one_latent(value, f"object_id:{key}")
|
||||||
|
if parsed is not None:
|
||||||
|
self.manual_object_latent[str(key)] = parsed
|
||||||
|
for key, value in agent_section.items():
|
||||||
|
parsed = self._parse_one_latent(value, f"agent_id:{key}")
|
||||||
|
if parsed is not None:
|
||||||
|
self.manual_agent_latent[str(key)] = parsed
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _parse_one_latent(value: dict, name: str) -> Optional[Dict[str, np.ndarray]]:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
print(f"[HBBC] invalid latent entry ({name}): expect dict.")
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
eps = _normalize_eps(value["latent_eps"])
|
||||||
|
c = _normalize_c(value["latent_c"])
|
||||||
|
return {"latent_eps": eps, "latent_c": c}
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[HBBC] invalid latent entry ({name}): {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _lookup_manual(self, object_id: Optional[str], agent_id: Optional[str]) -> Optional[Tuple[np.ndarray, np.ndarray]]:
|
||||||
|
if object_id is not None and object_id in self.manual_object_latent:
|
||||||
|
e = self.manual_object_latent[object_id]["latent_eps"]
|
||||||
|
c = self.manual_object_latent[object_id]["latent_c"]
|
||||||
|
return e, c
|
||||||
|
if agent_id is not None and agent_id in self.manual_agent_latent:
|
||||||
|
e = self.manual_agent_latent[agent_id]["latent_eps"]
|
||||||
|
c = self.manual_agent_latent[agent_id]["latent_c"]
|
||||||
|
return e, c
|
||||||
|
if self.manual_global_latent is not None:
|
||||||
|
return self.manual_global_latent
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_latent(self, vehicle_key: str, object_id: Optional[str], agent_id: Optional[str]) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
manual = self._lookup_manual(object_id=object_id, agent_id=agent_id)
|
||||||
|
if manual is not None:
|
||||||
|
return manual
|
||||||
|
if self.mode == "per_episode_reset":
|
||||||
|
if self._episode_latent is None:
|
||||||
|
self._episode_latent = _sample_latent(self.rng)
|
||||||
|
return self._episode_latent
|
||||||
|
if vehicle_key not in self.vehicle_latent:
|
||||||
|
self.vehicle_latent[vehicle_key] = _sample_latent(self.rng)
|
||||||
|
return self.vehicle_latent[vehicle_key]
|
||||||
|
|
||||||
|
|
||||||
|
class HBBCBackgroundController:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_path: str,
|
||||||
|
device: str = "cpu",
|
||||||
|
latent_mode: str = "per_vehicle_fixed",
|
||||||
|
latent_json_path: Optional[str] = None,
|
||||||
|
seed: int = 0,
|
||||||
|
dt: float = 0.1,
|
||||||
|
):
|
||||||
|
self.model = HBBCModelWrapper.get(model_path=model_path, device=device)
|
||||||
|
self.latent_mgr = HBBCLatentManager(mode=latent_mode, seed=seed, latent_json_path=latent_json_path)
|
||||||
|
self.dt = float(dt)
|
||||||
|
self.vehicle_state: Dict[str, VehicleStateCache] = {}
|
||||||
|
|
||||||
|
def reset_episode(self):
|
||||||
|
self.latent_mgr.reset_episode()
|
||||||
|
self.vehicle_state.clear()
|
||||||
|
|
||||||
|
def remove_vehicle(self, vehicle_key: str):
|
||||||
|
self.vehicle_state.pop(vehicle_key, None)
|
||||||
|
self.latent_mgr.vehicle_latent.pop(vehicle_key, None)
|
||||||
|
|
||||||
|
def _build_base_state(self, vehicle, vehicle_key: str) -> np.ndarray:
|
||||||
|
state = self.vehicle_state.get(vehicle_key)
|
||||||
|
if state is None:
|
||||||
|
state = VehicleStateCache()
|
||||||
|
self.vehicle_state[vehicle_key] = state
|
||||||
|
|
||||||
|
speed_km_h = float(getattr(vehicle, "speed_km_h", 0.0))
|
||||||
|
max_speed_km_h = float(getattr(vehicle, "max_speed_km_h", 120.0))
|
||||||
|
veh_vel = np.clip((speed_km_h + 1.0) / (max_speed_km_h + 1.0), 0.0, 1.0)
|
||||||
|
|
||||||
|
heading_theta = float(getattr(vehicle, "heading_theta", 0.0))
|
||||||
|
if state.last_heading_theta is None:
|
||||||
|
yaw_rate = 0.0
|
||||||
|
else:
|
||||||
|
yaw_rate = _wrap_to_pi(heading_theta - state.last_heading_theta) / self.dt
|
||||||
|
yaw_rate = float(np.clip(yaw_rate, -5.0, 5.0))
|
||||||
|
|
||||||
|
current_action = getattr(vehicle, "current_action", None)
|
||||||
|
if current_action is None:
|
||||||
|
last_action_0, last_action_1 = state.last_action
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
last_action_0, last_action_1 = float(current_action[0]), float(current_action[1])
|
||||||
|
except Exception:
|
||||||
|
last_action_0, last_action_1 = state.last_action
|
||||||
|
|
||||||
|
state.last_heading_theta = heading_theta
|
||||||
|
state.last_speed_km_h = speed_km_h
|
||||||
|
state.last_action = (last_action_0, last_action_1)
|
||||||
|
|
||||||
|
obs = np.array(
|
||||||
|
[
|
||||||
|
0.0,
|
||||||
|
0.0,
|
||||||
|
0.0,
|
||||||
|
veh_vel,
|
||||||
|
0.0,
|
||||||
|
yaw_rate * 0.5,
|
||||||
|
last_action_0,
|
||||||
|
last_action_1,
|
||||||
|
],
|
||||||
|
dtype=np.float32,
|
||||||
|
)
|
||||||
|
return obs
|
||||||
|
|
||||||
|
def build_obs(self, vehicle, vehicle_key: str, object_id: Optional[str], agent_id: Optional[str]) -> np.ndarray:
|
||||||
|
base = self._build_base_state(vehicle, vehicle_key=vehicle_key)
|
||||||
|
eps, c = self.latent_mgr.get_latent(vehicle_key=vehicle_key, object_id=object_id, agent_id=agent_id)
|
||||||
|
return np.concatenate([base, eps, c], axis=-1).astype(np.float32)
|
||||||
|
|
||||||
|
def infer_actions(self, batch: List[Tuple[str, object, Optional[str], Optional[str]]]) -> Dict[str, np.ndarray]:
|
||||||
|
if not batch:
|
||||||
|
return {}
|
||||||
|
obs_list = []
|
||||||
|
vehicle_ids = []
|
||||||
|
for vehicle_key, vehicle, object_id, agent_id in batch:
|
||||||
|
obs_list.append(self.build_obs(vehicle, vehicle_key=vehicle_key, object_id=object_id, agent_id=agent_id))
|
||||||
|
vehicle_ids.append(vehicle_key)
|
||||||
|
actions = self.model.act_batch(np.stack(obs_list, axis=0))
|
||||||
|
out = {}
|
||||||
|
for idx, key in enumerate(vehicle_ids):
|
||||||
|
out[key] = actions[idx].astype(np.float32)
|
||||||
|
return out
|
||||||
@@ -63,3 +63,31 @@ class InverseDynamics:
|
|||||||
norm_steering = np.clip(steering / self.max_steering, -1.0, 1.0)
|
norm_steering = np.clip(steering / self.max_steering, -1.0, 1.0)
|
||||||
|
|
||||||
return np.array([norm_steering, norm_acc]), {'raw_acc': acc, 'raw_steering': steering}
|
return np.array([norm_steering, norm_acc]), {'raw_acc': acc, 'raw_steering': steering}
|
||||||
|
|
||||||
|
def apply_action(self, current_state, action, dt=0.1):
|
||||||
|
"""
|
||||||
|
Forward dynamics: given current_state and action [steering, acc] in [-1, 1], return next_state.
|
||||||
|
State format: dict with position (x,y), heading, velocity (vx, vy).
|
||||||
|
"""
|
||||||
|
steering_norm, acc_norm = float(action[0]), float(action[1])
|
||||||
|
acc = acc_norm * self.max_acc
|
||||||
|
steering = steering_norm * self.max_steering
|
||||||
|
pos = np.array(current_state['position'][:2], dtype=np.float64)
|
||||||
|
heading = float(current_state['heading'])
|
||||||
|
vel = np.array(current_state['velocity'], dtype=np.float64)
|
||||||
|
v = np.linalg.norm(vel)
|
||||||
|
if v < 0.1:
|
||||||
|
v = 0.1
|
||||||
|
theta_dot = v * np.tan(steering) / self.wheelbase
|
||||||
|
v_next = v + acc * dt
|
||||||
|
v_next = max(0.0, v_next)
|
||||||
|
heading_next = heading + theta_dot * dt
|
||||||
|
heading_next = np.arctan2(np.sin(heading_next), np.cos(heading_next))
|
||||||
|
vx_next = v_next * np.cos(heading_next)
|
||||||
|
vy_next = v_next * np.sin(heading_next)
|
||||||
|
pos_next = pos + dt * np.array([vx_next, vy_next])
|
||||||
|
return {
|
||||||
|
'position': pos_next,
|
||||||
|
'heading': heading_next,
|
||||||
|
'velocity': np.array([vx_next, vy_next]),
|
||||||
|
}
|
||||||
|
|||||||
@@ -53,6 +53,13 @@ class MultiAgentScenarioEnv(ScenarioEnv):
|
|||||||
data_directory=None,
|
data_directory=None,
|
||||||
num_controlled_agents=3,
|
num_controlled_agents=3,
|
||||||
horizon=1000,
|
horizon=1000,
|
||||||
|
# HBBC background vehicle controls (optional)
|
||||||
|
enable_hbbc_background=False,
|
||||||
|
hbbc_model_path="models/hbbc/hbbc.pt",
|
||||||
|
hbbc_inference_device="cpu",
|
||||||
|
hbbc_latent_mode="per_vehicle_fixed",
|
||||||
|
hbbc_latent_json_path=None,
|
||||||
|
hbbc_dt=0.1,
|
||||||
))
|
))
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
|||||||
31
README.md
31
README.md
@@ -16,6 +16,7 @@ MAGAIL4AutoDrive/
|
|||||||
│ └── ...
|
│ └── ...
|
||||||
├── Env/ # 仿真环境封装 (MetaDrive Wrapper)
|
├── Env/ # 仿真环境封装 (MetaDrive Wrapper)
|
||||||
│ ├── bc_env.py # BCScenarioEnv,45 维观测(BC/MAGAIL 共用)
|
│ ├── bc_env.py # BCScenarioEnv,45 维观测(BC/MAGAIL 共用)
|
||||||
|
│ ├── bc_ego_replay_env.py # BCEgoReplayEnv,单智能体 BC 评估(仅 ego 受控)
|
||||||
│ ├── scenario_env.py # 多智能体基础场景环境
|
│ ├── scenario_env.py # 多智能体基础场景环境
|
||||||
│ ├── expert_replay_env.py # 专家轨迹回放环境(数据生成与回放)
|
│ ├── expert_replay_env.py # 专家轨迹回放环境(数据生成与回放)
|
||||||
│ ├── inverse_dynamics.py # 逆动力学模块 (轨迹 -> 动作)
|
│ ├── inverse_dynamics.py # 逆动力学模块 (轨迹 -> 动作)
|
||||||
@@ -78,26 +79,47 @@ python -m scenarionet.convert_waymo -d data/exp_converted --raw_data_path ./waym
|
|||||||
**4) 本项目:生成专家 pkl**
|
**4) 本项目:生成专家 pkl**
|
||||||
使用筛选后的场景目录,生成训练用 pkl 到 `data/training_data`:
|
使用筛选后的场景目录,生成训练用 pkl 到 `data/training_data`:
|
||||||
|
|
||||||
|
- **多智能体**(所有受控车轨迹,输出 `expert_data_{start_index}_{num_scenarios}.pkl`):
|
||||||
```bash
|
```bash
|
||||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||||
```
|
```
|
||||||
|
|
||||||
|
- **单智能体**(仅 ego 车轨迹,输出 `expert_data_ego_{start_index}_{num_scenarios}.pkl`,用于单智能体 BC):
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||||
|
```
|
||||||
|
|
||||||
## 核心工作流
|
## 核心工作流
|
||||||
|
|
||||||
### 1. 数据准备
|
### 1. 数据准备
|
||||||
使用 `scripts/generate_expert_data.py` 将 Waymo 数据转换为训练用 `.pkl`,输出到 `data/training_data/`。
|
使用 `scripts/generate_expert_data.py` 将 Waymo 数据转换为训练用 `.pkl`,输出到 `data/training_data/`。
|
||||||
|
|
||||||
|
- **多智能体**:
|
||||||
```bash
|
```bash
|
||||||
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||||
|
```
|
||||||
|
|
||||||
|
- **单智能体(仅 ego)**:
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||||
```
|
```
|
||||||
|
|
||||||
### 2. 行为克隆 (BC)
|
### 2. 行为克隆 (BC)
|
||||||
- **训练**:`python train_bc.py`(模型保存到 `models/bc/`,日志到 `logs/bc/`)
|
BC 支持两种模式:**多智能体**(默认,所有受控车共用同一策略)与 **单智能体**(仅 ego 车,评估时其他车按专家轨迹回放)。
|
||||||
|
|
||||||
|
- **多智能体训练**(模型保存到 `models/bc/`,日志到 `logs/bc/`):
|
||||||
|
```bash
|
||||||
|
python train_bc.py --expert_data_path data/training_data/expert_data_0_50.pkl --epochs 100
|
||||||
```
|
```
|
||||||
# 注意替换文件名
|
|
||||||
python train_bc.py --expert_data_path ./data/training/expert_data_0_50.pkl --epochs 100
|
- **单智能体训练**(使用 ego-only 数据,评估时仅 ego 受策略控制,其他车专家回放):
|
||||||
|
```bash
|
||||||
|
python train_bc.py --expert_data_path data/training_data/expert_data_ego_0_50.pkl --epochs 100 --single_agent
|
||||||
```
|
```
|
||||||
|
|
||||||
- **可视化**:`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt`
|
- **可视化**:`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt`
|
||||||
|
仅自车用策略、其他车回放(单智能体可视化):加 `--ego_only`,例如
|
||||||
|
`python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --ego_only --num_scenarios 1`
|
||||||
|
|
||||||
### 3. 多智能体对抗模仿学习 (MAGAIL)
|
### 3. 多智能体对抗模仿学习 (MAGAIL)
|
||||||
- **训练**:`python train_magail.py`(模型保存到 `models/magail/`,日志到 `logs/magail/`)
|
- **训练**:`python train_magail.py`(模型保存到 `models/magail/`,日志到 `logs/magail/`)
|
||||||
@@ -114,6 +136,7 @@ python train_bc.py --expert_data_path ./data/training/expert_data_0_50.pkl --epo
|
|||||||
|
|
||||||
### Env 模块
|
### Env 模块
|
||||||
- **Env/bc_env.py**:`BCScenarioEnv`,45 维观测(Ego 5 维 + 10 邻居×4 维),BC 与 MAGAIL 训练/评估共用
|
- **Env/bc_env.py**:`BCScenarioEnv`,45 维观测(Ego 5 维 + 10 邻居×4 维),BC 与 MAGAIL 训练/评估共用
|
||||||
|
- **Env/bc_ego_replay_env.py**:`BCEgoReplayEnv`,单智能体 BC 评估环境,仅 ego 受策略控制,其他车按专家轨迹回放
|
||||||
- **Env/scenario_env.py**:`MultiAgentScenarioEnv` 基类,Waymo 场景加载与步进
|
- **Env/scenario_env.py**:`MultiAgentScenarioEnv` 基类,Waymo 场景加载与步进
|
||||||
- **Env/expert_replay_env.py**:专家轨迹回放与逆动力学动作,供 `generate_expert_data.py` 与回放可视化
|
- **Env/expert_replay_env.py**:专家轨迹回放与逆动力学动作,供 `generate_expert_data.py` 与回放可视化
|
||||||
- **Env/inverse_dynamics.py**:轨迹 → 油门/转向动作
|
- **Env/inverse_dynamics.py**:轨迹 → 油门/转向动作
|
||||||
|
|||||||
2
algorithms/__init__.py
Normal file
2
algorithms/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
|||||||
|
"""Compatibility package for legacy HBBC checkpoints."""
|
||||||
|
|
||||||
18
algorithms/utils.py
Normal file
18
algorithms/utils.py
Normal file
@@ -0,0 +1,18 @@
|
|||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
class RunningMeanStd(object):
|
||||||
|
def __init__(self, epsilon=1e-4, shape=()):
|
||||||
|
self.mean = np.zeros(shape, np.float64)
|
||||||
|
self.var = np.ones(shape, np.float64)
|
||||||
|
self.count = epsilon
|
||||||
|
|
||||||
|
|
||||||
|
class Normalizer(RunningMeanStd):
|
||||||
|
def __init__(self, input_dim, epsilon=1e-4, clip_obs=10.0):
|
||||||
|
super().__init__(shape=input_dim)
|
||||||
|
self.epsilon = epsilon
|
||||||
|
self.clip_obs = clip_obs
|
||||||
|
|
||||||
|
def normalize(self, input):
|
||||||
|
return np.clip((input - self.mean) / np.sqrt(self.var + self.epsilon), -self.clip_obs, self.clip_obs)
|
||||||
@@ -10,7 +10,7 @@ import torch
|
|||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
|
||||||
def load_expert_pkl(expert_data_path, *, filter_terminal_last_step: bool = False):
|
def load_expert_pkl(expert_data_path, *, filter_terminal_last_step: bool = False, agent_id_filter=None):
|
||||||
"""从目录或单个 pkl 加载专家 (obs, acts),返回 concat 后的 obs_data, act_data。
|
"""从目录或单个 pkl 加载专家 (obs, acts),返回 concat 后的 obs_data, act_data。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -18,6 +18,8 @@ def load_expert_pkl(expert_data_path, *, filter_terminal_last_step: bool = False
|
|||||||
filter_terminal_last_step: If True, drop the last (obs, act) pair of each trajectory.
|
filter_terminal_last_step: If True, drop the last (obs, act) pair of each trajectory.
|
||||||
This approximates II's \"train only on non-terminal steps\" when the dataset doesn't
|
This approximates II's \"train only on non-terminal steps\" when the dataset doesn't
|
||||||
explicitly store dones.
|
explicitly store dones.
|
||||||
|
agent_id_filter: If not None, only load trajectories with traj[\"agent_id\"] == agent_id_filter
|
||||||
|
(e.g. \"default_agent\" for single-agent/ego-only).
|
||||||
"""
|
"""
|
||||||
if os.path.isdir(expert_data_path):
|
if os.path.isdir(expert_data_path):
|
||||||
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
pkl_files = glob.glob(os.path.join(expert_data_path, "*.pkl"))
|
||||||
@@ -36,6 +38,8 @@ def load_expert_pkl(expert_data_path, *, filter_terminal_last_step: bool = False
|
|||||||
data = pickle.load(f)
|
data = pickle.load(f)
|
||||||
if isinstance(data, list):
|
if isinstance(data, list):
|
||||||
for traj in data:
|
for traj in data:
|
||||||
|
if agent_id_filter is not None and traj.get("agent_id") != agent_id_filter:
|
||||||
|
continue
|
||||||
if "obs" in traj and "acts" in traj:
|
if "obs" in traj and "acts" in traj:
|
||||||
obs = traj["obs"]
|
obs = traj["obs"]
|
||||||
acts = traj["acts"]
|
acts = traj["acts"]
|
||||||
@@ -101,11 +105,12 @@ def get_expert_scenario_ids(expert_data_path, max_ids=10):
|
|||||||
|
|
||||||
|
|
||||||
class MAGAILExpertDataset(Dataset):
|
class MAGAILExpertDataset(Dataset):
|
||||||
def __init__(self, data_dir, transform=None, *, filter_terminal_last_step: bool = False):
|
def __init__(self, data_dir, transform=None, *, filter_terminal_last_step: bool = False, agent_id_filter=None):
|
||||||
"""
|
"""
|
||||||
Args:
|
Args:
|
||||||
data_dir (str): Directory containing .pkl files from generate_expert_data.py
|
data_dir (str): Directory containing .pkl files from generate_expert_data.py
|
||||||
transform (callable, optional): Optional transform to be applied on a sample.
|
transform (callable, optional): Optional transform to be applied on a sample.
|
||||||
|
agent_id_filter: If not None, only load trajectories with traj[\"agent_id\"] == agent_id_filter.
|
||||||
"""
|
"""
|
||||||
self.data_dir = data_dir
|
self.data_dir = data_dir
|
||||||
self.transform = transform
|
self.transform = transform
|
||||||
@@ -121,6 +126,8 @@ class MAGAILExpertDataset(Dataset):
|
|||||||
with open(pkl_file, "rb") as f:
|
with open(pkl_file, "rb") as f:
|
||||||
data = pickle.load(f)
|
data = pickle.load(f)
|
||||||
# data is a list of dicts: {'obs': (T, 45), 'acts': (T, 2), ...}
|
# data is a list of dicts: {'obs': (T, 45), 'acts': (T, 2), ...}
|
||||||
|
if agent_id_filter is not None:
|
||||||
|
data = [t for t in data if t.get("agent_id") == agent_id_filter]
|
||||||
self.trajectories.extend(data)
|
self.trajectories.extend(data)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error loading {pkl_file}: {e}")
|
print(f"Error loading {pkl_file}: {e}")
|
||||||
|
|||||||
439
docs/HBBC_Deploy_guied.md
Normal file
439
docs/HBBC_Deploy_guied.md
Normal file
@@ -0,0 +1,439 @@
|
|||||||
|
# HBBC 策略部署指南
|
||||||
|
|
||||||
|
本文档说明如何将 `weights/hbbc.pt` 部署到 MetaDrive 项目中的**背景车辆**上,作为车辆控制策略使用。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 0. 本仓库适配说明(MAGAIL4AutoDrive)
|
||||||
|
|
||||||
|
本仓库已落地一套可直接使用的 HBBC 背景车接入实现,核心代码:
|
||||||
|
|
||||||
|
- `Env/hbbc_actor_critic.py`:HBBC 所需 `ActorCritic` 最小推理网络
|
||||||
|
- `Env/hbbc_background_policy.py`:模型加载、18 维观测构建、latent 管理(含 JSON 覆盖)
|
||||||
|
- `Env/bc_env.py`:`BCScenarioEnv` 动态背景车 HBBC 接入(静态背景车保持不变)
|
||||||
|
- `Env/bc_ego_replay_env.py`:`BCEgoReplayEnv` 动态背景车 HBBC 接入(ego-only 评估兼容)
|
||||||
|
|
||||||
|
与原文档示例不同点:
|
||||||
|
|
||||||
|
1. 当前仓库 `BaseVehicle` 没有 `pos_buffer/rot_buffer/action_buffer`,因此 8 维 `base_state` 使用当前可得车辆状态重建;
|
||||||
|
2. 仅动态背景车使用 HBBC,静态背景车仍作为占位/邻居车辆;
|
||||||
|
3. 支持通过 JSON 手动指定场景中某些车辆的 latent(`object_id` / `agent_id` 双 key)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. 概述
|
||||||
|
|
||||||
|
### 1.1 HBBC 是什么
|
||||||
|
|
||||||
|
**HBBC**(Hierarchical Behavior-Based Controller)是一个低层驾驶策略网络,输入车辆状态和行为条件,输出连续控制动作 `[steering, acceleration]`,可直接用于 MetaDrive 的车辆控制。
|
||||||
|
|
||||||
|
### 1.2 依赖
|
||||||
|
|
||||||
|
- **PyTorch**
|
||||||
|
- **NumPy**
|
||||||
|
- **MetaDrive**(需包含 `BaseVehicle`、`BasePolicy` 等基础组件)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. 模型加载
|
||||||
|
|
||||||
|
### 2.1 模型架构
|
||||||
|
|
||||||
|
HBBC 对应 `ActorCritic` 网络,需按以下参数实例化:
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch
|
||||||
|
from algorithms.modules import ActorCritic # 或复制 actor_critic.py 到目标项目
|
||||||
|
|
||||||
|
hbbc = ActorCritic(
|
||||||
|
num_actor_obs=18,
|
||||||
|
num_critic_obs=18,
|
||||||
|
num_actions=2,
|
||||||
|
latent_c_dim=4, # 行为模式数
|
||||||
|
latent_eps_dim=6, # 风格向量维度
|
||||||
|
use_style_latent=True,
|
||||||
|
).to(device)
|
||||||
|
|
||||||
|
# 加载权重
|
||||||
|
checkpoint = torch.load("path/to/hbbc.pt", map_location=device, weights_only=False)
|
||||||
|
hbbc.load_state_dict(checkpoint['actor_critic'])
|
||||||
|
hbbc.eval()
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.2 推理接口
|
||||||
|
|
||||||
|
```python
|
||||||
|
with torch.no_grad():
|
||||||
|
actions = hbbc.act_inference(obs_tensor) # obs_tensor: (batch, 18), 输出: (batch, 2)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. 输入规格(18 维)
|
||||||
|
|
||||||
|
HBBC 的输入为 `hbbc_obs`,维度 18,由三部分拼接:
|
||||||
|
|
||||||
|
```
|
||||||
|
hbbc_obs = [base_state(8) | latent_eps(6) | latent_c(4)]
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.1 base_state(8 维)
|
||||||
|
|
||||||
|
从车辆对象构建,需按**精确顺序**拼接。实现如下(需配合 `relative_pos_local`、`rot_matrix_inv`、`clip` 等工具函数):
|
||||||
|
|
||||||
|
```python
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
def build_hbbc_base_state(vehicle):
|
||||||
|
"""
|
||||||
|
从 MetaDrive 车辆对象构建 HBBC 的 8 维 base_state。
|
||||||
|
要求 vehicle 具有: position, pos_buffer, rot_buffer, heading_buffer,
|
||||||
|
speed_km_h, max_speed_km_h, eps_step, acceleration, yaw_rate, action_buffer
|
||||||
|
"""
|
||||||
|
from metadrive.utils.math import clip # 或 np.clip
|
||||||
|
|
||||||
|
veh_pos = list(vehicle.position) + [0]
|
||||||
|
init_veh_rot = np.array([vehicle.rot_buffer[0][0], vehicle.rot_buffer[0][1], vehicle.rot_buffer[0][2]])
|
||||||
|
init_veh_pos = list(vehicle.pos_buffer[0]) + [0]
|
||||||
|
init_veh_heading = vehicle.heading_buffer[0]
|
||||||
|
|
||||||
|
# 局部位置(本实现中置 0)
|
||||||
|
veh_pos_local = relative_pos_local(init_veh_pos, veh_pos, init_veh_rot)[:2]
|
||||||
|
veh_pos_local[0] /= 10
|
||||||
|
veh_pos_local[1] /= 2
|
||||||
|
|
||||||
|
# 局部航向(本实现中置 0)
|
||||||
|
veh_heading = vehicle.heading
|
||||||
|
cross = np.cross(init_veh_heading, veh_heading)
|
||||||
|
dot = np.dot(init_veh_heading, veh_heading)
|
||||||
|
veh_heading_local = np.arctan2(cross, dot)
|
||||||
|
|
||||||
|
veh_vel = clip((vehicle.speed_km_h + 1) / (vehicle.max_speed_km_h + 1), 0.0, 1.0)
|
||||||
|
veh_acc = vehicle.acceleration / 5 if vehicle.eps_step > 1 else 0
|
||||||
|
yaw_rate = vehicle.yaw_rate
|
||||||
|
last_action_0 = vehicle.action_buffer[-1][0]
|
||||||
|
last_action_1 = vehicle.action_buffer[-1][1]
|
||||||
|
|
||||||
|
# 8 维,顺序固定
|
||||||
|
obs = np.concatenate((
|
||||||
|
veh_pos_local * 0, # 2 维,置 0
|
||||||
|
[veh_heading_local * 0], # 1 维,置 0
|
||||||
|
[veh_vel], # 1 维
|
||||||
|
[veh_acc * 0], # 1 维,置 0
|
||||||
|
[yaw_rate * 0.5], # 1 维
|
||||||
|
[last_action_0], [last_action_1] # 2 维
|
||||||
|
)).astype(np.float32)
|
||||||
|
return obs
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.2 latent_eps(6 维)
|
||||||
|
|
||||||
|
风格向量,需 **L2 归一化** 且在 `[-1, 1]` 内:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 随机采样(每个 episode 或每辆车可固定/随机)
|
||||||
|
latent_eps = np.random.randn(6).astype(np.float32)
|
||||||
|
latent_eps = latent_eps / (np.linalg.norm(latent_eps) + 1e-8)
|
||||||
|
latent_eps = np.clip(latent_eps, -1.0, 1.0)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.3 latent_c(4 维)
|
||||||
|
|
||||||
|
行为模式 one-hot,4 选 1:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 随机选一个模式 (0~3)
|
||||||
|
mode = np.random.randint(0, 4)
|
||||||
|
latent_c = np.zeros(4, dtype=np.float32)
|
||||||
|
latent_c[mode] = 1.0
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.4 完整观测拼接
|
||||||
|
|
||||||
|
```python
|
||||||
|
def build_hbbc_obs(vehicle, latent_eps, latent_c):
|
||||||
|
base = build_hbbc_base_state(vehicle)
|
||||||
|
return np.concatenate([base, latent_eps, latent_c], axis=-1) # shape: (18,)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. 必需工具函数
|
||||||
|
|
||||||
|
若目标项目无以下函数,需自行实现或从 styledrive 的 `envs/utils.py` 拷贝:
|
||||||
|
|
||||||
|
```python
|
||||||
|
def rot_matrix(t):
|
||||||
|
"""t: [roll, pitch, yaw], 返回 3x3 旋转矩阵"""
|
||||||
|
roll, pitch, yaw = t[0], t[1], t[2]
|
||||||
|
sr, cr = np.sin(roll), np.cos(roll)
|
||||||
|
sp, cp = np.sin(pitch), np.cos(pitch)
|
||||||
|
sy, cy = np.sin(yaw), np.cos(yaw)
|
||||||
|
r_roll = np.array([[1, 0, 0], [0, cr, -sr], [0, sr, cr]])
|
||||||
|
r_pitch = np.array([[cp, 0, sp], [0, 1, 0], [-sp, 0, cp]])
|
||||||
|
r_yaw = np.array([[cy, -sy, 0], [sy, cy, 0], [0, 0, 1]])
|
||||||
|
return np.dot(np.dot(r_yaw, r_pitch), r_roll)
|
||||||
|
|
||||||
|
def rot_matrix_inv(t):
|
||||||
|
return rot_matrix(t).T
|
||||||
|
|
||||||
|
def relative_pos_local(coord, coord_t, veh_rot):
|
||||||
|
"""将 coord_t 从世界坐标变换到以 coord 为原点、veh_rot 为姿态的局部坐标"""
|
||||||
|
r_pos_global = np.array(coord_t) - np.array(coord)
|
||||||
|
rot_mat_inv = rot_matrix_inv(veh_rot)
|
||||||
|
return rot_mat_inv @ r_pos_global
|
||||||
|
```
|
||||||
|
|
||||||
|
`clip` 可用 `np.clip` 或 `metadrive.utils.math.clip`。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. 车辆属性要求
|
||||||
|
|
||||||
|
使用 HBBC 的车辆需继承或兼容 MetaDrive 的 `BaseVehicle`,并具备:
|
||||||
|
|
||||||
|
| 属性 | 说明 |
|
||||||
|
|------|------|
|
||||||
|
| `position` | 当前位置 (x, y) 或 (x, y, z) |
|
||||||
|
| `heading` | 航向单位向量 |
|
||||||
|
| `heading_theta` | 航向角(弧度) |
|
||||||
|
| `pos_buffer` | `deque`,至少 1 个元素,`pos_buffer[0]` 为 episode 起始位姿 |
|
||||||
|
| `rot_buffer` | `deque`,`(roll, pitch, yaw)`,`rot_buffer[0]` 为起始姿态 |
|
||||||
|
| `heading_buffer` | `deque`,`heading_buffer[0]` 为起始航向 |
|
||||||
|
| `action_buffer` | `deque`,`action_buffer[-1]` 为上一时刻动作 `(steering, acc)` |
|
||||||
|
| `speed_km_h` | 当前速度 km/h |
|
||||||
|
| `max_speed_km_h` | 最大速度 km/h |
|
||||||
|
| `acceleration` | 当前加速度 |
|
||||||
|
| `yaw_rate` | 偏航角速度 (rad/s) |
|
||||||
|
| `eps_step` | 本 episode 的步数 |
|
||||||
|
| `last_heading_theta` | 上一帧航向角(用于 yaw_rate) |
|
||||||
|
|
||||||
|
`BaseVehicle` 在 `before_step` 中会更新 `pos_buffer`、`rot_buffer`、`heading_buffer`、`action_buffer`,只要在配置中设置 `veh_obs_len >= 1`(建议 3–10)即可。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. 输出动作格式
|
||||||
|
|
||||||
|
HBBC 输出 2 维连续动作,与 MetaDrive 动作空间一致:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# actions: (2,) 或 (batch, 2)
|
||||||
|
# actions[0]: steering ∈ [-1, 1]
|
||||||
|
# actions[1]: acceleration ∈ [-1, 1],正=油门,负=刹车
|
||||||
|
```
|
||||||
|
|
||||||
|
环境会在 `_preprocess_actions` 中做限幅与平滑,无需在策略内再次裁剪。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. 部署为 MetaDrive 策略(背景车)
|
||||||
|
|
||||||
|
### 7.1 自定义 Policy
|
||||||
|
|
||||||
|
实现一个继承 `BasePolicy` 的策略,在 `act` 中调用 HBBC:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from metadrive.policy.base_policy import BasePolicy
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
class HBBCPolicy(BasePolicy):
|
||||||
|
def __init__(self, control_object, random_seed=None, hbbc_path="weights/hbbc.pt", device="cpu"):
|
||||||
|
super().__init__(control_object, random_seed)
|
||||||
|
self.device = torch.device(device)
|
||||||
|
self.hbbc = self._load_hbbc(hbbc_path)
|
||||||
|
self.latent_eps = None
|
||||||
|
self.latent_c = None
|
||||||
|
self._resample_latent()
|
||||||
|
|
||||||
|
def _load_hbbc(self, path):
|
||||||
|
from algorithms.modules import ActorCritic # 根据实际路径调整
|
||||||
|
model = ActorCritic(
|
||||||
|
num_actor_obs=18, num_critic_obs=18, num_actions=2,
|
||||||
|
latent_c_dim=4, latent_eps_dim=6, use_style_latent=True
|
||||||
|
).to(self.device)
|
||||||
|
ckpt = torch.load(path, map_location=self.device, weights_only=False)
|
||||||
|
model.load_state_dict(ckpt['actor_critic'])
|
||||||
|
model.eval()
|
||||||
|
return model
|
||||||
|
|
||||||
|
def _resample_latent(self):
|
||||||
|
self.latent_eps = np.random.randn(6).astype(np.float32)
|
||||||
|
self.latent_eps = self.latent_eps / (np.linalg.norm(self.latent_eps) + 1e-8)
|
||||||
|
self.latent_eps = np.clip(self.latent_eps, -1.0, 1.0)
|
||||||
|
mode = np.random.randint(0, 4)
|
||||||
|
self.latent_c = np.zeros(4, dtype=np.float32)
|
||||||
|
self.latent_c[mode] = 1.0
|
||||||
|
|
||||||
|
def act(self, agent_id=None):
|
||||||
|
vehicle = self.control_object
|
||||||
|
base_state = build_hbbc_base_state(vehicle)
|
||||||
|
obs = np.concatenate([base_state, self.latent_eps, self.latent_c], axis=-1)
|
||||||
|
obs_t = torch.tensor(obs, dtype=torch.float32, device=self.device).unsqueeze(0)
|
||||||
|
with torch.no_grad():
|
||||||
|
actions = self.hbbc.act_inference(obs_t).cpu().numpy().squeeze()
|
||||||
|
self.action_info["action"] = actions.tolist()
|
||||||
|
return [float(actions[0]), float(actions[1])]
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
super().reset()
|
||||||
|
self._resample_latent()
|
||||||
|
```
|
||||||
|
|
||||||
|
### 7.2 配置背景车使用 HBBC
|
||||||
|
|
||||||
|
在环境配置中为背景车辆指定 `HBBCPolicy`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
config = {
|
||||||
|
# ...
|
||||||
|
"agent_configs": {
|
||||||
|
"agent0": {
|
||||||
|
"policy": HBBCPolicy,
|
||||||
|
"policy_kwargs": {"hbbc_path": "path/to/hbbc.pt", "device": "cuda:0"},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
# 若使用 traffic 的 policy 配置方式,则需在 traffic 管理逻辑中
|
||||||
|
# 将部分或全部背景车的 policy 替换为 HBBCPolicy
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
若背景车由 TrafficManager 等模块统一管理,需在该模块的 policy 选择逻辑中加入对 `HBBCPolicy` 的分配。
|
||||||
|
|
||||||
|
### 7.3 与 TrafficManager 集成
|
||||||
|
|
||||||
|
若背景车由 `PGTrafficManager` 等生成,需在添加策略时改为使用 `HBBCPolicy`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 原代码通常为:
|
||||||
|
# self.add_policy(random_v.id, IDMPolicy, random_v, self.generate_seed())
|
||||||
|
|
||||||
|
# 改为:
|
||||||
|
from your_policy_module import HBBCPolicy
|
||||||
|
self.add_policy(random_v.id, HBBCPolicy, random_v, self.generate_seed(),
|
||||||
|
hbbc_path="path/to/hbbc.pt", device="cuda:0")
|
||||||
|
```
|
||||||
|
|
||||||
|
`add_policy` 的额外参数会传给 Policy 的 `__init__`。若接口不支持传参,可修改 `HBBCPolicy` 从全局配置读取路径,或使用自定义 TrafficManager 子类。
|
||||||
|
|
||||||
|
**注意**:HBBC 在 styledrive 中基于 scenario 轨迹训练,不包含路由逻辑。背景车若需要沿车道/路线行驶,可能需:
|
||||||
|
- 在项目中为 HBBC 车辆配置 `navigation`,或
|
||||||
|
- 仅对部分背景车使用 HBBC(如混合 IDM + HBBC),或
|
||||||
|
- 在目标项目中验证 HBBC 在开放道路上的表现后决定是否全量使用。
|
||||||
|
|
||||||
|
### 7.4 注意事项
|
||||||
|
|
||||||
|
1. **latent 生命周期**:可为每辆车在 spawn 时采样一次,或在每个 episode reset 时重采样。
|
||||||
|
2. **首帧 action_buffer**:首步 `action_buffer[-1]` 通常为 `(0, 0)`,由 `BaseVehicle` 初始化保证。
|
||||||
|
3. **同步更新 buffer**:车辆必须在每步调用 `before_step` 之类接口,更新 `pos_buffer`、`action_buffer` 等,否则观测会错位。
|
||||||
|
4. **veh_obs_len**:车辆配置中设置 `veh_obs_len >= 3`(建议 10),确保 buffer 长度足够。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 8. ActorCritic 网络定义(可移植)
|
||||||
|
|
||||||
|
若目标项目无法导入 styledrive 的 `algorithms`,可把以下简化版 `ActorCritic` 放到本项目中单独使用:
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
def get_activation(name):
|
||||||
|
return getattr(nn, name)()
|
||||||
|
|
||||||
|
class ActorCritic(nn.Module):
|
||||||
|
def __init__(self, num_actor_obs=18, num_critic_obs=18, num_actions=2,
|
||||||
|
latent_c_dim=4, latent_eps_dim=6, use_style_latent=True,
|
||||||
|
actor_hidden_dims=[512, 256, 128], activation='elu'):
|
||||||
|
super().__init__()
|
||||||
|
act_fn = getattr(nn, activation.upper())()
|
||||||
|
self.latent_c_dim = latent_c_dim
|
||||||
|
self.latent_eps_dim = latent_eps_dim
|
||||||
|
self.use_style_latent = use_style_latent
|
||||||
|
|
||||||
|
layers = []
|
||||||
|
layers.append(nn.Linear(num_actor_obs, actor_hidden_dims[0]))
|
||||||
|
layers.append(act_fn)
|
||||||
|
for i in range(len(actor_hidden_dims) - 1):
|
||||||
|
layers.append(nn.Linear(actor_hidden_dims[i], actor_hidden_dims[i + 1]))
|
||||||
|
layers.append(act_fn)
|
||||||
|
self.actor_trunk = nn.Sequential(*layers)
|
||||||
|
self.actor_head = nn.Linear(actor_hidden_dims[-1], num_actions)
|
||||||
|
|
||||||
|
if use_style_latent:
|
||||||
|
style_layers = [nn.Linear(latent_eps_dim, 512), act_fn,
|
||||||
|
nn.Linear(512, 256), act_fn, nn.Linear(256, 128), act_fn]
|
||||||
|
self.style_trunk = nn.Sequential(*style_layers)
|
||||||
|
self.style_head = nn.Linear(128, latent_eps_dim)
|
||||||
|
self.style_activation = torch.tanh
|
||||||
|
|
||||||
|
def act_inference(self, observations):
|
||||||
|
if self.use_style_latent:
|
||||||
|
obs = observations[..., :-(self.latent_c_dim + self.latent_eps_dim)]
|
||||||
|
eps = observations[..., -self.latent_c_dim - self.latent_eps_dim:-self.latent_c_dim]
|
||||||
|
c = observations[..., -self.latent_c_dim:]
|
||||||
|
eps = self.style_activation(self.style_head(self.style_trunk(eps)))
|
||||||
|
observations = torch.cat([obs, eps, c], dim=-1)
|
||||||
|
embedding = self.actor_trunk(observations)
|
||||||
|
return self.actor_head(embedding)
|
||||||
|
```
|
||||||
|
|
||||||
|
加载与调用方式与前面一致。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 9. 简要检查清单
|
||||||
|
|
||||||
|
- [ ] 正确加载 `hbbc.pt` 的 `actor_critic` 权重
|
||||||
|
- [ ] `build_hbbc_base_state` 输出 8 维,顺序与文档一致
|
||||||
|
- [ ] `latent_eps` 6 维、L2 归一化
|
||||||
|
- [ ] `latent_c` 4 维 one-hot
|
||||||
|
- [ ] 车辆具备 `pos_buffer`、`rot_buffer`、`heading_buffer`、`action_buffer` 等属性
|
||||||
|
- [ ] 策略返回 `[steering, acceleration]`,范围 [-1, 1]
|
||||||
|
- [ ] 每步更新上述 buffer,保证观测连续
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 10. 本仓库配置项与 JSON 示例
|
||||||
|
|
||||||
|
可通过环境配置控制 HBBC 背景车行为:
|
||||||
|
|
||||||
|
- `enable_hbbc_background`:是否启用动态背景车 HBBC(`True/False`)
|
||||||
|
- `hbbc_model_path`:模型路径(默认 `models/hbbc/hbbc.pt`)
|
||||||
|
- `hbbc_inference_device`:推理设备(如 `cpu` / `cuda:0`)
|
||||||
|
- `hbbc_latent_mode`:`per_vehicle_fixed` 或 `per_episode_reset`
|
||||||
|
- `hbbc_latent_json_path`:可选,手动 latent JSON 路径
|
||||||
|
|
||||||
|
`hbbc_latent_json_path` 内容格式(优先按 `object_id` 匹配,失败回退 `agent_id`):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"global": {
|
||||||
|
"latent_eps": [0.35, -0.12, 0.28, 0.46, -0.22, 0.18],
|
||||||
|
"latent_c": [0, 0, 1, 0]
|
||||||
|
},
|
||||||
|
"object_id": {
|
||||||
|
"12345": {
|
||||||
|
"latent_eps": [0.2, -0.1, 0.3, 0.4, -0.2, 0.1],
|
||||||
|
"latent_c": [0, 1, 0, 0]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agent_id": {
|
||||||
|
"controlled_abcde": {
|
||||||
|
"latent_eps": [0.5, 0.1, -0.1, 0.2, -0.3, 0.4],
|
||||||
|
"latent_c": [1, 0, 0, 0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
匹配优先级为:`object_id` > `agent_id` > `global` > 随机采样。
|
||||||
|
`latent_eps` 会做 L2 归一化,`latent_c` 会强制 one-hot;非法输入会告警并回退随机采样。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 11. 参考来源
|
||||||
|
|
||||||
|
- 策略与观测:`envs/ad_hbbc_gym.py` 中的 `ADObservation.vehicle_state`
|
||||||
|
- 模型:`algorithms/modules/actor_critic.py` 中 `ActorCritic`
|
||||||
|
- 工具:`envs/utils.py` 中的 `relative_pos_local`、`rot_matrix`、`rot_matrix_inv`
|
||||||
18
docs/examples/hbbc_latent_example.json
Normal file
18
docs/examples/hbbc_latent_example.json
Normal file
@@ -0,0 +1,18 @@
|
|||||||
|
{
|
||||||
|
"global": {
|
||||||
|
"latent_eps": [0.35, -0.12, 0.28, 0.46, -0.22, 0.18],
|
||||||
|
"latent_c": [0, 1,1, 0]
|
||||||
|
},
|
||||||
|
"object_id": {
|
||||||
|
"12345": {
|
||||||
|
"latent_eps": [0.2, -0.1, 0.3, 0.4, -0.2, 0.1],
|
||||||
|
"latent_c": [0, 1, 0, 0]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agent_id": {
|
||||||
|
"controlled_abcde": {
|
||||||
|
"latent_eps": [0.5, 0.1, -0.1, 0.2, -0.3, 0.4],
|
||||||
|
"latent_c": [1, 0, 0, 0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Binary file not shown.
@@ -16,9 +16,19 @@
|
|||||||
|
|
||||||
| 脚本 | 用途 | 用法示例 |
|
| 脚本 | 用途 | 用法示例 |
|
||||||
|------|------|----------|
|
|------|------|----------|
|
||||||
| [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` |
|
| [generate_expert_data.py](generate_expert_data.py) | 从 Waymo 数据生成专家 (obs, act) 的 pkl | 见下方 |
|
||||||
|
|
||||||
**常用参数**:`--data_dir`(默认 `data/exp_filtered`)、`--output_dir`(默认 `data/training_data`)、`--start_index`、`--num_scenarios`。
|
**多智能体**(输出 `expert_data_{start_index}_{num_scenarios}.pkl`):
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0
|
||||||
|
```
|
||||||
|
|
||||||
|
**单智能体**(仅采集 ego 车轨迹,输出 `expert_data_ego_{start_index}_{num_scenarios}.pkl`,用于单智能体 BC):
|
||||||
|
```bash
|
||||||
|
python scripts/generate_expert_data.py --data_dir data/exp_filtered --output_dir data/training_data --num_scenarios 100 --start_index 0 --ego_only
|
||||||
|
```
|
||||||
|
|
||||||
|
**常用参数**:`--data_dir`(默认 `data/exp_filtered`)、`--output_dir`(默认 `data/training_data`)、`--start_index`、`--num_scenarios`、`--ego_only`(仅保存 default_agent 轨迹,输出使用 `expert_data_ego_*.pkl` 前缀)。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -40,13 +50,32 @@ python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios
|
|||||||
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
|
||||||
```
|
```
|
||||||
|
- **policy + 仅自车策略、其他车回放**(BC 单智能体模型):加 `--ego_only`,自车由策略控制,其余车辆按专家轨迹回放。
|
||||||
|
```bash
|
||||||
|
python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1 --ego_only
|
||||||
|
```
|
||||||
|
|
||||||
|
- **policy + HBBC 动态背景车**(仅动态背景车启用,静态背景车保持原样):
|
||||||
|
```bash
|
||||||
|
python scripts/visualize.py policy \
|
||||||
|
--policy_type bc \
|
||||||
|
--model_path models/bc/policy_best.pt \
|
||||||
|
--data_dir data/exp_filtered \
|
||||||
|
--num_scenarios 1 \
|
||||||
|
--ego_only \
|
||||||
|
--enable_hbbc_background \
|
||||||
|
--hbbc_model_path models/hbbc/hbbc.pt \
|
||||||
|
--hbbc_inference_device cpu \
|
||||||
|
--hbbc_latent_mode per_vehicle_fixed \
|
||||||
|
--hbbc_latent_json_path docs/examples/hbbc_latent_example.json
|
||||||
|
```
|
||||||
|
|
||||||
- **trajectory**(专家轨迹 matplotlib 俯视图动画):
|
- **trajectory**(专家轨迹 matplotlib 俯视图动画):
|
||||||
```bash
|
```bash
|
||||||
python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_idx 0
|
python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_idx 0
|
||||||
```
|
```
|
||||||
|
|
||||||
**公共参数**:`--data_dir`(默认 `data/exp_filtered`)、`--start_index`、`--num_scenarios`、`--horizon`。policy 模式另有 `--policy_type`(auto/bc/magail)、`--model_path`、`--deterministic`(仅 MAGAIL)。
|
**公共参数**:`--data_dir`(默认 `data/exp_filtered`)、`--start_index`、`--num_scenarios`、`--horizon`。policy 模式另有 `--policy_type`(auto/bc/magail)、`--model_path`、`--deterministic`(仅 MAGAIL)、`--ego_only`(仅 BC:自车用策略,其他车专家回放)、`--enable_hbbc_background`、`--hbbc_model_path`、`--hbbc_inference_device`、`--hbbc_latent_mode`、`--hbbc_latent_json_path`。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -70,7 +99,7 @@ python scripts/visualize.py trajectory --data_dir data/exp_filtered --scenario_i
|
|||||||
|
|
||||||
## 与训练流程的对应关系
|
## 与训练流程的对应关系
|
||||||
|
|
||||||
1. **数据准备**:`generate_expert_data.py` → 输出到 `data/training_data/*.pkl`
|
1. **数据准备**:`generate_expert_data.py` → 输出到 `data/training_data/*.pkl`(多智能体 `expert_data_*.pkl`,单智能体 `expert_data_ego_*.pkl`)
|
||||||
2. **BC 训练**:根目录 `train_bc.py` → 模型保存到 `models/bc/`,日志到 `logs/bc/`
|
2. **BC 训练**:根目录 `train_bc.py` → 模型保存到 `models/bc/`,日志到 `logs/bc/`。单智能体模式加 `--single_agent` 并指定 ego-only 的 pkl。
|
||||||
3. **MAGAIL 训练**:根目录 `train_magail.py` → 模型保存到 `models/magail/`,日志到 `logs/magail/`
|
3. **MAGAIL 训练**:根目录 `train_magail.py` → 模型保存到 `models/magail/`,日志到 `logs/magail/`
|
||||||
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
4. **可视化**:`scripts/visualize.py`(子命令 replay / policy / trajectory)→ 数据目录默认 `data/exp_filtered`
|
||||||
|
|||||||
@@ -102,6 +102,8 @@ def generate_data(args):
|
|||||||
|
|
||||||
# Post-process episode data
|
# Post-process episode data
|
||||||
for agent_id, data in episode_data.items():
|
for agent_id, data in episode_data.items():
|
||||||
|
if args.ego_only and agent_id != "default_agent":
|
||||||
|
continue
|
||||||
if len(data['obs']) > 10: # Minimum length filter
|
if len(data['obs']) > 10: # Minimum length filter
|
||||||
expert_trajectories.append({
|
expert_trajectories.append({
|
||||||
'obs': np.array(data['obs']),
|
'obs': np.array(data['obs']),
|
||||||
@@ -120,9 +122,14 @@ def generate_data(args):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
# Save data
|
# Save data
|
||||||
|
if args.ego_only:
|
||||||
|
output_file = os.path.join(args.output_dir, f"expert_data_ego_{args.start_index}_{args.num_scenarios}.pkl")
|
||||||
|
else:
|
||||||
output_file = os.path.join(args.output_dir, f"expert_data_{args.start_index}_{args.num_scenarios}.pkl")
|
output_file = os.path.join(args.output_dir, f"expert_data_{args.start_index}_{args.num_scenarios}.pkl")
|
||||||
os.makedirs(args.output_dir, exist_ok=True)
|
os.makedirs(args.output_dir, exist_ok=True)
|
||||||
|
|
||||||
|
if args.ego_only:
|
||||||
|
print("Ego-only mode: saved trajectories are SDC (default_agent) only.")
|
||||||
print(f"Saving {len(expert_trajectories)} trajectories to {output_file}")
|
print(f"Saving {len(expert_trajectories)} trajectories to {output_file}")
|
||||||
with open(output_file, 'wb') as f:
|
with open(output_file, 'wb') as f:
|
||||||
pickle.dump(expert_trajectories, f)
|
pickle.dump(expert_trajectories, f)
|
||||||
@@ -157,6 +164,6 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument("--output_dir", type=str, default="data/training_data", help="Output directory")
|
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("--start_index", type=int, default=0)
|
||||||
parser.add_argument("--num_scenarios", type=int, default=10)
|
parser.add_argument("--num_scenarios", type=int, default=10)
|
||||||
|
parser.add_argument("--ego_only", action="store_true", help="Only collect and save ego (default_agent) trajectories; output uses expert_data_ego_*.pkl prefix")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
generate_data(args)
|
generate_data(args)
|
||||||
|
|||||||
@@ -111,36 +111,49 @@ def _resolve_model_path(model_path, policy_type):
|
|||||||
|
|
||||||
def _run_policy(args):
|
def _run_policy(args):
|
||||||
from Env.bc_env import BCScenarioEnv
|
from Env.bc_env import BCScenarioEnv
|
||||||
|
from Env.bc_ego_replay_env import BCEgoReplayEnv
|
||||||
from metadrive.engine.engine_utils import close_engine
|
from metadrive.engine.engine_utils import close_engine
|
||||||
|
|
||||||
policy_type = (args.policy_type or "auto").lower()
|
policy_type = (args.policy_type or "auto").lower()
|
||||||
if policy_type == "auto":
|
if policy_type == "auto":
|
||||||
policy_type = "bc" if args.model_path.endswith(".pt") else "magail"
|
policy_type = "bc" if args.model_path.endswith(".pt") else "magail"
|
||||||
|
ego_only = getattr(args, "ego_only", False)
|
||||||
|
if ego_only and policy_type != "bc":
|
||||||
|
print("[WARN] --ego_only is supported for BC policy only; MAGAIL will run in multi-agent mode.")
|
||||||
|
|
||||||
data_dir = _resolve_data_dir(args.data_dir)
|
data_dir = _resolve_data_dir(args.data_dir)
|
||||||
data_path = os.path.abspath(data_dir)
|
data_path = os.path.abspath(data_dir)
|
||||||
env_config = {
|
env_config = {
|
||||||
"data_directory": data_path,
|
"data_directory": data_path,
|
||||||
"is_multi_agent": True,
|
"is_multi_agent": True,
|
||||||
"num_controlled_agents": 3,
|
"num_controlled_agents": 100 if ego_only else 3,
|
||||||
"horizon": args.horizon,
|
"horizon": args.horizon,
|
||||||
"use_render": True,
|
"use_render": True,
|
||||||
"sequential_seed": True,
|
"sequential_seed": True,
|
||||||
"start_scenario_index": args.start_index,
|
"start_scenario_index": args.start_index,
|
||||||
"num_scenarios": args.num_scenarios,
|
"num_scenarios": args.num_scenarios,
|
||||||
"log_level": 40,
|
"log_level": 40,
|
||||||
|
"enable_hbbc_background": bool(getattr(args, "enable_hbbc_background", False)),
|
||||||
|
"hbbc_model_path": getattr(args, "hbbc_model_path", "models/hbbc/hbbc.pt"),
|
||||||
|
"hbbc_inference_device": getattr(args, "hbbc_inference_device", "cpu"),
|
||||||
|
"hbbc_latent_mode": getattr(args, "hbbc_latent_mode", "per_vehicle_fixed"),
|
||||||
|
"hbbc_latent_json_path": getattr(args, "hbbc_latent_json_path", None),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if ego_only and policy_type == "bc":
|
||||||
|
print("Initializing BCEgoReplayEnv (ego-only: policy on self, others replayed)...")
|
||||||
|
else:
|
||||||
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
print(f"Initializing BCScenarioEnv (policy_type={policy_type})...")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
env = BCScenarioEnv(env_config, agent2policy={})
|
env = BCEgoReplayEnv(config=env_config) if (ego_only and policy_type == "bc") else BCScenarioEnv(env_config, agent2policy={})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error init env: {e}. Trying to close lingering engine...")
|
print(f"Error init env: {e}. Trying to close lingering engine...")
|
||||||
try:
|
try:
|
||||||
close_engine()
|
close_engine()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
env = BCScenarioEnv(env_config, agent2policy={})
|
env = BCEgoReplayEnv(config=env_config) if (ego_only and policy_type == "bc") else BCScenarioEnv(env_config, agent2policy={})
|
||||||
|
|
||||||
state_dim = 45
|
state_dim = 45
|
||||||
action_dim = 2
|
action_dim = 2
|
||||||
@@ -156,7 +169,11 @@ def _run_policy(args):
|
|||||||
hidden_units=(256, 256),
|
hidden_units=(256, 256),
|
||||||
hidden_activation=torch.nn.Tanh(),
|
hidden_activation=torch.nn.Tanh(),
|
||||||
).to(device)
|
).to(device)
|
||||||
policy.load_state_dict(torch.load(model_path, map_location=device))
|
try:
|
||||||
|
state = torch.load(model_path, map_location=device, weights_only=True)
|
||||||
|
except TypeError:
|
||||||
|
state = torch.load(model_path, map_location=device)
|
||||||
|
policy.load_state_dict(state)
|
||||||
policy.eval()
|
policy.eval()
|
||||||
else:
|
else:
|
||||||
from train_magail import Actor
|
from train_magail import Actor
|
||||||
@@ -178,7 +195,16 @@ def _run_policy(args):
|
|||||||
pass
|
pass
|
||||||
continue
|
continue
|
||||||
|
|
||||||
print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {env.num_controlled_in_scenario}")
|
n_total = getattr(env, "num_controlled_in_scenario", len(obs_dict))
|
||||||
|
mode_note = " (ego only, others replayed)" if (ego_only and policy_type == "bc") else ""
|
||||||
|
if ego_only and policy_type == "bc" and bool(env_config.get("enable_hbbc_background", False)):
|
||||||
|
mode_note = " (ego only, dynamic background via HBBC)"
|
||||||
|
print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {n_total}{mode_note}")
|
||||||
|
if ego_only and policy_type == "bc" and len(obs_dict) == 1:
|
||||||
|
if bool(env_config.get("enable_hbbc_background", False)):
|
||||||
|
print(" [Ego control: policy injected — dynamic background vehicles use HBBC; static background stays static.]")
|
||||||
|
else:
|
||||||
|
print(" [Ego control: policy injected — ego uses model output each step; other vehicles expert replay.]")
|
||||||
if len(obs_dict) == 0:
|
if len(obs_dict) == 0:
|
||||||
print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.")
|
print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.")
|
||||||
continue
|
continue
|
||||||
@@ -370,6 +396,12 @@ def main():
|
|||||||
pp.add_argument("--policy_type", type=str, default="auto", choices=["auto", "bc", "magail"])
|
pp.add_argument("--policy_type", type=str, default="auto", choices=["auto", "bc", "magail"])
|
||||||
pp.add_argument("--model_path", type=str, default="models/bc/policy_best.pt")
|
pp.add_argument("--model_path", type=str, default="models/bc/policy_best.pt")
|
||||||
pp.add_argument("--deterministic", action="store_true", help="MAGAIL: use mean action")
|
pp.add_argument("--deterministic", action="store_true", help="MAGAIL: use mean action")
|
||||||
|
pp.add_argument("--ego_only", action="store_true", help="BC only: inject policy into ego only; other vehicles use expert replay")
|
||||||
|
pp.add_argument("--enable_hbbc_background", action="store_true", help="Enable HBBC policy for dynamic background vehicles")
|
||||||
|
pp.add_argument("--hbbc_model_path", type=str, default="models/hbbc/hbbc.pt")
|
||||||
|
pp.add_argument("--hbbc_inference_device", type=str, default="cpu")
|
||||||
|
pp.add_argument("--hbbc_latent_mode", type=str, default="per_vehicle_fixed", choices=["per_vehicle_fixed", "per_episode_reset"])
|
||||||
|
pp.add_argument("--hbbc_latent_json_path", type=str, default=None, help="Optional JSON for per-vehicle latent override")
|
||||||
|
|
||||||
# trajectory
|
# trajectory
|
||||||
pt = subparsers.add_parser("trajectory", help="2D matplotlib animation of expert trajectories")
|
pt = subparsers.add_parser("trajectory", help="2D matplotlib animation of expert trajectories")
|
||||||
|
|||||||
25
train_bc.py
25
train_bc.py
@@ -15,19 +15,20 @@ 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 Env.bc_ego_replay_env import BCEgoReplayEnv
|
||||||
from dataset.loader import load_expert_pkl, get_expert_scenario_ids
|
from dataset.loader import load_expert_pkl, get_expert_scenario_ids
|
||||||
|
|
||||||
|
|
||||||
def evaluate_policy(policy, args, device):
|
def evaluate_policy(policy, args, device):
|
||||||
"""在 BCScenarioEnv 中评估策略:仅使用专家数据中出现过的 scenario_id,保证 eval 有受控车。
|
"""在 BCScenarioEnv(多智能体)或 BCEgoReplayEnv(单智能体)中评估策略。
|
||||||
输出与 replay 对齐:agents (current)=reset 时受控车数,total in scenario=该场景受控轨迹总数(car_birth_info_list 长度)。"""
|
仅使用专家数据中出现过的 scenario_id。单智能体模式下仅 ego 受策略控制,其他车专家回放。"""
|
||||||
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):
|
||||||
data_dir = os.path.join(waymo_data_dir, "exp_converted")
|
data_dir = os.path.join(waymo_data_dir, "exp_converted")
|
||||||
if not os.path.exists(data_dir):
|
if not os.path.exists(data_dir):
|
||||||
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, 0.0, 0.0
|
||||||
|
|
||||||
scenario_ids = get_expert_scenario_ids(args.expert_data_path, max_ids=5)
|
scenario_ids = get_expert_scenario_ids(args.expert_data_path, max_ids=5)
|
||||||
if not scenario_ids:
|
if not scenario_ids:
|
||||||
@@ -38,6 +39,7 @@ def evaluate_policy(policy, args, device):
|
|||||||
total_steps = []
|
total_steps = []
|
||||||
collision_episodes = 0
|
collision_episodes = 0
|
||||||
horizon = 200
|
horizon = 200
|
||||||
|
single_agent = getattr(args, "single_agent", False)
|
||||||
|
|
||||||
for idx, scenario_id in enumerate(scenario_ids):
|
for idx, scenario_id in enumerate(scenario_ids):
|
||||||
env_config = {
|
env_config = {
|
||||||
@@ -49,7 +51,11 @@ def evaluate_policy(policy, args, device):
|
|||||||
"horizon": horizon,
|
"horizon": horizon,
|
||||||
"start_scenario_index": scenario_id,
|
"start_scenario_index": scenario_id,
|
||||||
"num_scenarios": 1,
|
"num_scenarios": 1,
|
||||||
|
"log_level": 50,
|
||||||
}
|
}
|
||||||
|
if single_agent:
|
||||||
|
env = BCEgoReplayEnv(config=env_config)
|
||||||
|
else:
|
||||||
env = BCScenarioEnv(env_config, agent2policy=None)
|
env = BCScenarioEnv(env_config, agent2policy=None)
|
||||||
try:
|
try:
|
||||||
obs_dict = env.reset(seed=scenario_id)
|
obs_dict = env.reset(seed=scenario_id)
|
||||||
@@ -59,7 +65,7 @@ def evaluate_policy(policy, args, device):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
n_controlled = len(env.controlled_agents)
|
n_controlled = len(env.controlled_agents)
|
||||||
n_total_in_scenario = getattr(env, "num_controlled_in_scenario", n_controlled)
|
n_total_in_scenario = getattr(env, "num_controlled_in_scenario", n_controlled) if not single_agent else 1
|
||||||
if n_controlled == 0:
|
if n_controlled == 0:
|
||||||
print(
|
print(
|
||||||
f" Eval Episode {idx} (scenario {scenario_id}): 0 controlled agents (total in scenario: {n_total_in_scenario}), skip."
|
f" Eval Episode {idx} (scenario {scenario_id}): 0 controlled agents (total in scenario: {n_total_in_scenario}), skip."
|
||||||
@@ -95,9 +101,9 @@ def evaluate_policy(policy, args, device):
|
|||||||
total_steps.append(step_count)
|
total_steps.append(step_count)
|
||||||
if had_near_collision:
|
if had_near_collision:
|
||||||
collision_episodes += 1
|
collision_episodes += 1
|
||||||
|
mode_str = "single-agent (ego)" if single_agent else f"agents (current): {n_controlled}, total in scenario: {n_total_in_scenario}"
|
||||||
print(
|
print(
|
||||||
f" Eval Episode {idx} (scenario {scenario_id}): Total Reward {episode_reward:.2f}, steps {step_count}, "
|
f" Eval Episode {idx} (scenario {scenario_id}): Total Reward {episode_reward:.2f}, steps {step_count}, {mode_str}"
|
||||||
f"agents (current): {n_controlled}, total in scenario: {n_total_in_scenario}"
|
|
||||||
)
|
)
|
||||||
env.close()
|
env.close()
|
||||||
|
|
||||||
@@ -124,9 +130,11 @@ def main(args):
|
|||||||
print(f"TensorBoard logging to: {log_dir}")
|
print(f"TensorBoard logging to: {log_dir}")
|
||||||
os.makedirs(args.save_dir, exist_ok=True)
|
os.makedirs(args.save_dir, exist_ok=True)
|
||||||
|
|
||||||
|
agent_id_filter = "default_agent" if getattr(args, "single_agent", False) else None
|
||||||
obs_data, act_data = load_expert_pkl(
|
obs_data, act_data = load_expert_pkl(
|
||||||
args.expert_data_path,
|
args.expert_data_path,
|
||||||
filter_terminal_last_step=args.filter_terminal_last_step,
|
filter_terminal_last_step=args.filter_terminal_last_step,
|
||||||
|
agent_id_filter=agent_id_filter,
|
||||||
)
|
)
|
||||||
obs_tensor = torch.FloatTensor(obs_data)
|
obs_tensor = torch.FloatTensor(obs_data)
|
||||||
act_tensor = torch.FloatTensor(act_data)
|
act_tensor = torch.FloatTensor(act_data)
|
||||||
@@ -195,5 +203,10 @@ if __name__ == "__main__":
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Drop the last (obs, act) pair of each trajectory to approximate training on non-terminal steps (II-style).",
|
help="Drop the last (obs, act) pair of each trajectory to approximate training on non-terminal steps (II-style).",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--single_agent",
|
||||||
|
action="store_true",
|
||||||
|
help="Use single-agent (ego) expert data and evaluation; load only default_agent trajectories and evaluate with BCEgoReplayEnv.",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
main(args)
|
main(args)
|
||||||
|
|||||||
Reference in New Issue
Block a user