diff --git a/plugins/mtgs/server/main.py b/plugins/mtgs/server/main.py index 195fcc46..77870699 100644 --- a/plugins/mtgs/server/main.py +++ b/plugins/mtgs/server/main.py @@ -13,7 +13,7 @@ import logging from concurrent import futures from pathlib import Path -from typing import Callable +from typing import Callable, Optional from alpasim_grpc.v0.sensorsim_pb2_grpc import add_SensorsimServiceServicer_to_server from alpasim_mtgs.server.servicer import MTGSSensorsimService @@ -42,9 +42,62 @@ DATASET_NAME_MAPPING = { "nuplan_test": "navtest", "nuplan_mini": "navtest", + "nuplan_private": "private", } +def _build_token_to_asset_folder(configs_dir: Path) -> dict: + """Read MTGS config YAMLs to build a central_token → road_block_name mapping. + + Each YAML lists all central_tokens sharing one rendered asset folder + (road_block_name = central_log + '-' + central_tokens[0]). This mapping + lets the MTGS server resolve any token to its shared asset folder at + runtime, regardless of the trajdata cache state. + """ + import yaml + + class _SafeLoader(yaml.SafeLoader): + pass + + _SafeLoader.add_multi_constructor( + "tag:yaml.org,2002:python/object", + lambda loader, tag, node: loader.construct_mapping(node, deep=True), + ) + _SafeLoader.add_multi_constructor( + "tag:yaml.org,2002:python/tuple", + lambda loader, tag, node: loader.construct_sequence(node, deep=True), + ) + + mapping: dict = {} + if not configs_dir.exists(): + logger.warning("MTGS configs dir not found: %s", configs_dir) + return mapping + + for yaml_file in configs_dir.glob("*.yaml"): + try: + cfg = yaml.load(yaml_file.read_text(), Loader=_SafeLoader) + if not isinstance(cfg, dict): + continue + central_log = cfg.get("central_log", "") + central_tokens = cfg.get("central_tokens", []) + if not central_tokens: + continue + road_block_name = cfg.get("road_block_name", "") + if not road_block_name and central_log: + road_block_name = f"{central_log}-{central_tokens[0]}" + if not road_block_name: + continue + for token in central_tokens: + mapping[str(token)] = road_block_name + except Exception as exc: + logger.warning("Failed to load %s: %s", yaml_file.name, exc) + + logger.info( + "Built asset-folder mapping for %d tokens from %s", len(mapping), configs_dir + ) + return mapping + + def parse_args(arg_list: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser( description="MTGS Sensorsim Service Server", @@ -98,6 +151,17 @@ def create_get_scene_function( mtgs_asset_base_path = str(Path(asset_base_path_config) / mapped_name / "assets") logger.info(f"MTGS asset path: {mtgs_asset_base_path}") + configs_dir = Path(asset_base_path_config) / mapped_name / "configs" + token_to_asset_folder = _build_token_to_asset_folder(configs_dir) + + def _asset_folder_resolver(scene) -> str: + token = scene.name.rsplit("-", 1)[-1] + return token_to_asset_folder.get(str(token), scene.name) + + asset_folder_resolver: Optional[Callable] = ( + _asset_folder_resolver if token_to_asset_folder else None + ) + params = trajdata_provider_config_to_params(trajdata_config) logger.info("Creating UnifiedDataset from config") dataset = UnifiedDataset(**params) @@ -127,6 +191,7 @@ def get_scene(scene_id: str) -> TrajdataDataSource: vector_map_params=dataset.vector_map_params, smooth_trajectories=user_config.smooth_trajectories, asset_base_path=mtgs_asset_base_path, + asset_folder_resolver=asset_folder_resolver, ) scene_cache[scene_id] = data_source diff --git a/plugins/mtgs/tests/test_mtgs_plugin.py b/plugins/mtgs/tests/test_mtgs_plugin.py index ae378cd5..8dd2d0ce 100644 --- a/plugins/mtgs/tests/test_mtgs_plugin.py +++ b/plugins/mtgs/tests/test_mtgs_plugin.py @@ -138,6 +138,21 @@ def test_mtgs_scene_loader_requires_asset_base_path(monkeypatch): mtgs_main.create_get_scene_function(_mtgs_user_config()) +def test_build_token_to_asset_folder_infers_road_block_name(tmp_path): + from alpasim_mtgs.server import main as mtgs_main + + (tmp_path / "scene.yaml").write_text( + "central_log: log-a\n" "central_tokens:\n" " - token-1\n" " - token-2\n" + ) + + mapping = mtgs_main._build_token_to_asset_folder(tmp_path) + + assert mapping == { + "token-1": "log-a-token-1", + "token-2": "log-a-token-1", + } + + def test_mtgs_scene_loader_uses_public_trajdata_dataset_api(monkeypatch): from alpasim_mtgs.server import main as mtgs_main diff --git a/src/runtime/alpasim_runtime/scene_loader.py b/src/runtime/alpasim_runtime/scene_loader.py index 929c5648..6a70bded 100644 --- a/src/runtime/alpasim_runtime/scene_loader.py +++ b/src/runtime/alpasim_runtime/scene_loader.py @@ -120,24 +120,23 @@ def _tuple_constructor(loader, tag_suffix, node): for yaml_file in yaml_files: try: cfg = yaml.load(yaml_file.read_text(), Loader=_SafeLoader) - central_log = ( - cfg.get("central_log", "") - if isinstance(cfg, dict) - else getattr(cfg, "central_log", "") - ) - central_tokens = ( - cfg.get("central_tokens", []) - if isinstance(cfg, dict) - else getattr(cfg, "central_tokens", []) - ) + central_log = cfg.get("central_log", "") + central_tokens = cfg.get("central_tokens", []) if not central_log or not central_tokens: logger.warning( "%s missing central_log or central_tokens, skipping", yaml_file.name ) continue + road_block_name = ( + cfg.get("road_block_name", "") or f"{central_log}-{central_tokens[0]}" + ) for token in central_tokens: configs_by_log[central_log].append( - {"central_token": token, "logfile": central_log} + { + "central_token": token, + "logfile": central_log, + "asset_folder": road_block_name, + } ) except Exception as exc: logger.warning("Failed to load %s: %s", yaml_file.name, exc) diff --git a/src/trajdata/src/trajdata/dataset_specific/nuplan/nuplan_utils.py b/src/trajdata/src/trajdata/dataset_specific/nuplan/nuplan_utils.py index 68698be8..2150d144 100644 --- a/src/trajdata/src/trajdata/dataset_specific/nuplan/nuplan_utils.py +++ b/src/trajdata/src/trajdata/dataset_specific/nuplan/nuplan_utils.py @@ -225,17 +225,18 @@ def _load_scenes_from_central_tokens(self) -> List[Dict[str, str]]: # Create scene name using the central token format. scene_name = f"{logfile_name}-{central_token_hex}" - scenes.append( - { - "name": scene_name, - "location": _NUPLAN_SQL_MAP_FRIENDLY_NAMES_DICT.get( - location, location - ), - "num_timesteps": num_timesteps, - "start_idx": start_idx, - "end_idx": end_idx, - } - ) + scene_dict: Dict[str, Any] = { + "name": scene_name, + "location": _NUPLAN_SQL_MAP_FRIENDLY_NAMES_DICT.get( + location, location + ), + "num_timesteps": num_timesteps, + "start_idx": start_idx, + "end_idx": end_idx, + } + if "asset_folder" in config: + scene_dict["asset_folder"] = config["asset_folder"] + scenes.append(scene_dict) self.close_db()