diff --git a/Env/scenario_env.py b/Env/scenario_env.py index 6f1f531..2ee4a49 100644 --- a/Env/scenario_env.py +++ b/Env/scenario_env.py @@ -64,6 +64,11 @@ class MultiAgentScenarioEnv(ScenarioEnv): self.round = 0 super().__init__(config) + @property + def num_controlled_in_scenario(self) -> int: + """整个场景中受控车轨迹总数(car_birth_info_list 长度),会在不同 show_time 陆续 spawn。""" + return len(getattr(self, "car_birth_info_list", [])) + def reset(self, seed: Union[None, int] = None): self.round = 0 if self.logger is None: diff --git a/scripts/visualize.py b/scripts/visualize.py index 5a3e532..d7a9e47 100644 --- a/scripts/visualize.py +++ b/scripts/visualize.py @@ -55,7 +55,7 @@ def _run_replay(args): print(f"Error resetting scenario {i}: {e}") continue - print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}") + print(f"Scenario loaded. Controlled agents (current): {len(env.controlled_agents)}, total in scenario: {env.num_controlled_in_scenario}") for step in range(args.horizon): obs, rewards, dones, infos = env.step(None) @@ -176,7 +176,7 @@ def _run_policy(args): pass continue - print(f"Scenario loaded. Controlled agents: {len(obs_dict)}") + print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {env.num_controlled_in_scenario}") if len(obs_dict) == 0: print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.") continue