已合并
【poc】【feature】Support practice entry point for model adapter plugin #158
rookie_hongchuan创建于 2月27日
【poc】【feature】Support practice entry point for model adapter plugin #158
已合并
从已删除 :2026_Q1/support_practice_entry_point合入到Ascend/msmodelslimmaster
共 5 个文件变更+131-8
| @@ -384,7 +384,19 @@ class NaiveQuantizationApplication: | |||
| 384 | quant_type=quant_type, | 384 | quant_type=quant_type, |
| 385 | config_path=config_path | 385 | config_path=config_path |
| 386 | ) | 386 | ) |
| 387 | - get_logger().info(f"Get best practice {practice_config.metadata.config_id} success.") | 387 | + if config_path is not None: |
| 388 | + config_url = str(config_path) | ||
| 389 | + else: | ||
| 390 | + model_pedigree = model_adapter.get_model_pedigree() | ||
| 391 | + config_url = self.practice_manager.get_config_url( | ||
| 392 | + model_pedigree, practice_config.metadata.config_id | ||
| 393 | + ) | ||
| 394 | + if config_url is not None: | ||
| 395 | + get_logger().info( | ||
| 396 | + f"Get best practice {practice_config.metadata.config_id} success, config: {config_url}" | ||
| 397 | + ) | ||
| 398 | + else: | ||
| 399 | + get_logger().info(f"Get best practice {practice_config.metadata.config_id} success.") | ||
| 388 | 400 | ||
| 389 | get_logger().info(f"===========QUANTIZE MODEL===========") | 401 | get_logger().info(f"===========QUANTIZE MODEL===========") |
| 390 | self.quant_service.quantize( | 402 | self.quant_service.quantize( |
| @@ -19,12 +19,16 @@ See the Mulan PSL v2 for more details. | |||
| 19 | ------------------------------------------------------------------------- | 19 | ------------------------------------------------------------------------- |
| 20 | """ | 20 | """ |
| 21 | from abc import ABC, abstractmethod | 21 | from abc import ABC, abstractmethod |
| 22 | -from typing import Generator | 22 | +from typing import Generator, Optional |
| 23 | 23 | ||
| 24 | from msmodelslim.core.practice import PracticeConfig | 24 | from msmodelslim.core.practice import PracticeConfig |
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | class PracticeManagerInfra(ABC): | 27 | class PracticeManagerInfra(ABC): |
| 28 | + def get_config_url(self, model_pedigree: str, config_id: str) -> Optional[str]: | ||
| 29 | + """Return the URL/location of the config; in different scenarios url may be a file path, an HTTP URL, etc.""" | ||
| 30 | + return None | ||
| 31 | + | ||
| 28 | 32 | ||
| 29 | def __contains__(self, model_pedigree) -> bool: | 33 | def __contains__(self, model_pedigree) -> bool: |
| 30 | """Check if model pedigree is supported""" | 34 | """Check if model pedigree is supported""" |
| @@ -26,6 +26,7 @@ from msmodelslim.core.quant_service.proxy import QuantServiceProxy, QuantService | |||
| 26 | from msmodelslim.cli.utils import parse_device_string | 26 | from msmodelslim.cli.utils import parse_device_string |
| 27 | from msmodelslim.infra.file_dataset_loader import FileDatasetLoader | 27 | from msmodelslim.infra.file_dataset_loader import FileDatasetLoader |
| 28 | from msmodelslim.infra.dataset_loader.vlm_dataset_loader import VLMDatasetLoader | 28 | from msmodelslim.infra.dataset_loader.vlm_dataset_loader import VLMDatasetLoader |
| 29 | +from msmodelslim.infra.plugin_practice_dirs import discover_plugin_practice_dirs | ||
| 29 | from msmodelslim.infra.yaml_practice_manager import YamlPracticeManager | 30 | from msmodelslim.infra.yaml_practice_manager import YamlPracticeManager |
| 30 | from msmodelslim.model import PluginModelFactory | 31 | from msmodelslim.model import PluginModelFactory |
| 31 | from msmodelslim.utils.config import msmodelslim_config | 32 | from msmodelslim.utils.config import msmodelslim_config |
| @@ -50,9 +51,11 @@ def main(args): | |||
| 50 | config_dir = get_practice_dir() | 51 | config_dir = get_practice_dir() |
| 51 | custom_practice_dir = msmodelslim_config.env_vars.custom_practice_repo | 52 | custom_practice_dir = msmodelslim_config.env_vars.custom_practice_repo |
| 52 | custom_practice_path = Path(custom_practice_dir) if custom_practice_dir else None | 53 | custom_practice_path = Path(custom_practice_dir) if custom_practice_dir else None |
| 54 | + plugin_dirs = discover_plugin_practice_dirs() | ||
| 53 | practice_manager = YamlPracticeManager( | 55 | practice_manager = YamlPracticeManager( |
| 54 | official_config_dir=config_dir, | 56 | official_config_dir=config_dir, |
| 55 | - custom_config_dir=custom_practice_path | 57 | + custom_config_dir=custom_practice_path, |
| 58 | + third_party_config_dirs=plugin_dirs if plugin_dirs else None, | ||
| 56 | ) | 59 | ) |
| 57 | dataset_dir = get_dataset_dir() | 60 | dataset_dir = get_dataset_dir() |
| 58 | dataset_loader = FileDatasetLoader(dataset_dir) | 61 | dataset_loader = FileDatasetLoader(dataset_dir) |
| @@ -0,0 +1,53 @@ | |||
| 1 | +# -*- coding: UTF-8 -*- | ||
| 2 | +""" | ||
| 3 | +Discover practice directories from plugins via entry point msmodelslim.naive_quantization.practice_dirs. | ||
| 4 | + | ||
| 5 | +Each entry point is a callable (e.g. module:function) that returns a Path or a list of Paths | ||
| 6 | +pointing to practice root directories (same layout as lab_practice: pedigree subdirs with YAMLs). | ||
| 7 | +""" | ||
| 8 | +import sys | ||
| 9 | +from importlib.metadata import entry_points | ||
| 10 | +from pathlib import Path | ||
| 11 | +from typing import List | ||
| 12 | + | ||
| 13 | +from msmodelslim.utils.logging import get_logger | ||
| 14 | +from msmodelslim.utils.security import get_valid_read_path | ||
| 15 | + | ||
| 16 | +PRACTICE_DIRS_ENTRY_POINT = "msmodelslim.naive_quantization.practice_dirs" | ||
| 17 | + | ||
| 18 | + | ||
| 19 | +def discover_plugin_practice_dirs() -> List[Path]: | ||
| 20 | + """ | ||
| 21 | + Discover all plugin-provided practice root directories via entry point. | ||
| 22 | + Returns a list of valid, existing directory Paths (read-only). | ||
| 23 | + """ | ||
| 24 | + if sys.version_info >= (3, 10): | ||
| 25 | + eps = entry_points().select(group=PRACTICE_DIRS_ENTRY_POINT) | ||
| 26 | + else: | ||
| 27 | + eps = entry_points().get(PRACTICE_DIRS_ENTRY_POINT, []) | ||
| 28 | + | ||
| 29 | + result: List[Path] = [] | ||
| 30 | + for ep in eps: | ||
| 31 | + try: | ||
| 32 | + callable_obj = ep.load() | ||
| 33 | + value = callable_obj() | ||
| 34 | + if value is None: | ||
| 35 | + continue | ||
| 36 | + paths = [value] if isinstance(value, (Path, str)) else value | ||
| 37 | + for p in paths: | ||
| 38 | + path = Path(p) if isinstance(p, str) else p | ||
| 39 | + if not path.exists() or not path.is_dir(): | ||
| 40 | + continue | ||
| 41 | + get_valid_read_path(str(path), is_dir=True) | ||
| 42 | + result.append(path.resolve()) | ||
| 43 | + except Exception as e: # noqa: B902 | ||
| 44 | + get_logger().warning( | ||
| 45 | + "Failed to load practice dir from plugin %s: %s", ep.name, e | ||
| 46 | + ) | ||
| 47 | + if result: | ||
| 48 | + get_logger().info( | ||
| 49 | + "Discovered %d plugin practice dir(s): %s", | ||
| 50 | + len(result), | ||
| 51 | + [str(p) for p in result], | ||
| 52 | + ) | ||
| 53 | + return result | ||
| @@ -19,7 +19,7 @@ See the Mulan PSL v2 for more details. | |||
| 19 | ------------------------------------------------------------------------- | 19 | ------------------------------------------------------------------------- |
| 20 | """ | 20 | """ |
| 21 | from pathlib import Path | 21 | from pathlib import Path |
| 22 | -from typing import Dict, Generator, Optional | 22 | +from typing import Dict, Generator, List, Optional |
| 23 | 23 | ||
| 24 | from msmodelslim.app.auto_tuning import PracticeManagerInfra as atpm | 24 | from msmodelslim.app.auto_tuning import PracticeManagerInfra as atpm |
| 25 | from msmodelslim.app.naive_quantization import PracticeManagerInfra as nqpm | 25 | from msmodelslim.app.naive_quantization import PracticeManagerInfra as nqpm |
| @@ -33,7 +33,12 @@ class YamlPracticeManager( | |||
| 33 | nqpm, | 33 | nqpm, |
| 34 | atpm, | 34 | atpm, |
| 35 | ): | 35 | ): |
| 36 | - def __init__(self, official_config_dir: Path, custom_config_dir: Optional[Path] = None): | 36 | + def __init__( |
| 37 | + self, | ||
| 38 | + official_config_dir: Path, | ||
| 39 | + custom_config_dir: Optional[Path] = None, | ||
| 40 | + third_party_config_dirs: Optional[List[Path]] = None, | ||
| 41 | + ): | ||
| 37 | get_valid_read_path(str(official_config_dir), is_dir=True) | 42 | get_valid_read_path(str(official_config_dir), is_dir=True) |
| 38 | self.official_config_dir = official_config_dir | 43 | self.official_config_dir = official_config_dir |
| 39 | if custom_config_dir is not None: | 44 | if custom_config_dir is not None: |
| @@ -53,19 +58,44 @@ class YamlPracticeManager( | |||
| 53 | if model_type_dir.is_dir() | 58 | if model_type_dir.is_dir() |
| 54 | } if self.custom_config_dir else {} | 59 | } if self.custom_config_dir else {} |
| 55 | 60 | ||
| 61 | + # One dict per third-party root: pedigree -> YamlDatabase (read-only) | ||
| 62 | + self._plugin_database_maps: List[Dict[str, YamlDatabase]] = [] | ||
| 63 | + for plugin_root in (third_party_config_dirs or []): | ||
| 64 | + if not plugin_root.exists() or not plugin_root.is_dir(): | ||
| 65 | + continue | ||
| 66 | + try: | ||
| 67 | + get_valid_read_path(str(plugin_root), is_dir=True) | ||
| 68 | + db_map = { | ||
| 69 | + d.name: YamlDatabase(d, read_only=True) | ||
| 70 | + for d in plugin_root.iterdir() | ||
| 71 | + if d.is_dir() | ||
| 72 | + } | ||
| 73 | + if db_map: | ||
| 74 | + self._plugin_database_maps.append(db_map) | ||
| 75 | + except Exception: # noqa: S110 | ||
| 76 | + continue | ||
| 77 | + | ||
| 56 | def __contains__(self, model_pedigree: str) -> bool: | 78 | def __contains__(self, model_pedigree: str) -> bool: |
| 57 | model_pedigree = model_pedigree.lower() | 79 | model_pedigree = model_pedigree.lower() |
| 58 | - return model_pedigree in self.custom_databases or model_pedigree in self.official_databases | 80 | + if model_pedigree in self.custom_databases or model_pedigree in self.official_databases: |
| 81 | + return True | ||
| 82 | + return any(model_pedigree in m for m in self._plugin_database_maps) | ||
| 59 | 83 | ||
| 60 | def get_config_by_id(self, model_pedigree: str, config_id: str) -> PracticeConfig: | 84 | def get_config_by_id(self, model_pedigree: str, config_id: str) -> PracticeConfig: |
| 61 | model_pedigree = model_pedigree.lower() | 85 | model_pedigree = model_pedigree.lower() |
| 86 | + value = None | ||
| 62 | if model_pedigree in self.custom_databases and config_id in self.custom_databases[model_pedigree]: | 87 | if model_pedigree in self.custom_databases and config_id in self.custom_databases[model_pedigree]: |
| 63 | value = self.custom_databases[model_pedigree][config_id] | 88 | value = self.custom_databases[model_pedigree][config_id] |
| 64 | elif model_pedigree in self.official_databases and config_id in self.official_databases[model_pedigree]: | 89 | elif model_pedigree in self.official_databases and config_id in self.official_databases[model_pedigree]: |
| 65 | value = self.official_databases[model_pedigree][config_id] | 90 | value = self.official_databases[model_pedigree][config_id] |
| 66 | else: | 91 | else: |
| 67 | - raise UnsupportedError(f"Practice {config_id} of ModelType {model_pedigree} not found", | 92 | + for db_map in self._plugin_database_maps: |
| 68 | - action='Please check the practice id and model type') | 93 | + if model_pedigree in db_map and config_id in db_map[model_pedigree]: |
| 94 | + value = db_map[model_pedigree][config_id] | ||
| 95 | + break | ||
| 96 | + if value is None: | ||
| 97 | + raise UnsupportedError(f"Practice {config_id} of ModelType {model_pedigree} not found", | ||
| 98 | + action='Please check the practice id and model type') | ||
| 69 | 99 | ||
| 70 | quant_config = PracticeConfig.model_validate(value) | 100 | quant_config = PracticeConfig.model_validate(value) |
| 71 | 101 | ||
| @@ -74,11 +104,32 @@ class YamlPracticeManager( | |||
| 74 | action='Please make sure the practice is not tampered') | 104 | action='Please make sure the practice is not tampered') |
| 75 | return quant_config | 105 | return quant_config |
| 76 | 106 | ||
| 107 | + def get_config_url(self, model_pedigree: str, config_id: str) -> Optional[str]: | ||
| 108 | + """Return the URL/location of the config (same lookup order as get_config_by_id). In file-based use, url is the YAML path.""" | ||
| 109 | + path = self._get_config_path(model_pedigree, config_id) | ||
| 110 | + return str(path) if path is not None else None | ||
| 111 | + | ||
| 112 | + def _get_config_path(self, model_pedigree: str, config_id: str) -> Optional[Path]: | ||
| 113 | + """Return the full path to the YAML file for the given config (same lookup order as get_config_by_id).""" | ||
| 114 | + model_pedigree = model_pedigree.lower() | ||
| 115 | + if model_pedigree in self.custom_databases and config_id in self.custom_databases[model_pedigree]: | ||
| 116 | + return self.custom_databases[model_pedigree].config_dir / f"{config_id}.yaml" | ||
| 117 | + if model_pedigree in self.official_databases and config_id in self.official_databases[model_pedigree]: | ||
| 118 | + return self.official_databases[model_pedigree].config_dir / f"{config_id}.yaml" | ||
| 119 | + for db_map in self._plugin_database_maps: | ||
| 120 | + if model_pedigree in db_map and config_id in db_map[model_pedigree]: | ||
| 121 | + return db_map[model_pedigree].config_dir / f"{config_id}.yaml" | ||
| 122 | + return None | ||
| 123 | + | ||
| 77 | def iter_config(self, model_pedigree) -> Generator[PracticeConfig, None, None]: | 124 | def iter_config(self, model_pedigree) -> Generator[PracticeConfig, None, None]: |
| 78 | tasks = [] | 125 | tasks = [] |
| 79 | if model_pedigree in self.custom_databases: | 126 | if model_pedigree in self.custom_databases: |
| 80 | for value in self.custom_databases[model_pedigree].values(): | 127 | for value in self.custom_databases[model_pedigree].values(): |
| 81 | tasks.append(PracticeConfig.model_validate(value)) | 128 | tasks.append(PracticeConfig.model_validate(value)) |
| 129 | + for db_map in self._plugin_database_maps: | ||
| 130 | + if model_pedigree in db_map: | ||
| 131 | + for value in db_map[model_pedigree].values(): | ||
| 132 | + tasks.append(PracticeConfig.model_validate(value)) | ||
| 82 | if model_pedigree in self.official_databases: | 133 | if model_pedigree in self.official_databases: |
| 83 | for value in self.official_databases[model_pedigree].values(): | 134 | for value in self.official_databases[model_pedigree].values(): |
| 84 | tasks.append(PracticeConfig.model_validate(value)) | 135 | tasks.append(PracticeConfig.model_validate(value)) |