新增修改

This commit is contained in:
2026-02-04 21:07:09 +08:00
parent 95cc78d940
commit ceb6648a31
2 changed files with 7 additions and 2 deletions

View File

@@ -64,6 +64,11 @@ class MultiAgentScenarioEnv(ScenarioEnv):
self.round = 0 self.round = 0
super().__init__(config) super().__init__(config)
@property
def num_controlled_in_scenario(self) -> int:
"""整个场景中受控车轨迹总数car_birth_info_list 长度),会在不同 show_time 陆续 spawn。"""
return len(getattr(self, "car_birth_info_list", []))
def reset(self, seed: Union[None, int] = None): def reset(self, seed: Union[None, int] = None):
self.round = 0 self.round = 0
if self.logger is None: if self.logger is None:

View File

@@ -55,7 +55,7 @@ def _run_replay(args):
print(f"Error resetting scenario {i}: {e}") print(f"Error resetting scenario {i}: {e}")
continue continue
print(f"Scenario loaded. Controlled agents: {len(env.controlled_agents)}") print(f"Scenario loaded. Controlled agents (current): {len(env.controlled_agents)}, total in scenario: {env.num_controlled_in_scenario}")
for step in range(args.horizon): for step in range(args.horizon):
obs, rewards, dones, infos = env.step(None) obs, rewards, dones, infos = env.step(None)
@@ -176,7 +176,7 @@ def _run_policy(args):
pass pass
continue continue
print(f"Scenario loaded. Controlled agents: {len(obs_dict)}") print(f"Scenario loaded. Controlled agents (current): {len(obs_dict)}, total in scenario: {env.num_controlled_in_scenario}")
if len(obs_dict) == 0: 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