From ceb6648a3108f2dd8c122b5565d83008e875c5b2 Mon Sep 17 00:00:00 2001 From: huangfu <3045324663@qq.com> Date: Wed, 4 Feb 2026 21:07:09 +0800 Subject: [PATCH] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E4=BF=AE=E6=94=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Env/scenario_env.py | 5 +++++ scripts/visualize.py | 4 ++-- 2 files changed, 7 insertions(+), 2 deletions(-) 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