已合并
【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
已合并
rookie_hongchuan创建于 2月27日
从已删除 :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_path385 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"""
21from abc import ABC, abstractmethod21from abc import ABC, abstractmethod
22-from typing import Generator22+from typing import Generator, Optional
23 23 
24from msmodelslim.core.practice import PracticeConfig24from msmodelslim.core.practice import PracticeConfig
25 25 
26 26 
27class PracticeManagerInfra(ABC):27class 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 @abstractmethod32 @abstractmethod
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
26from msmodelslim.cli.utils import parse_device_string26from msmodelslim.cli.utils import parse_device_string
27from msmodelslim.infra.file_dataset_loader import FileDatasetLoader27from msmodelslim.infra.file_dataset_loader import FileDatasetLoader
28from msmodelslim.infra.dataset_loader.vlm_dataset_loader import VLMDatasetLoader28from msmodelslim.infra.dataset_loader.vlm_dataset_loader import VLMDatasetLoader
29+from msmodelslim.infra.plugin_practice_dirs import discover_plugin_practice_dirs
29from msmodelslim.infra.yaml_practice_manager import YamlPracticeManager30from msmodelslim.infra.yaml_practice_manager import YamlPracticeManager
30from msmodelslim.model import PluginModelFactory31from msmodelslim.model import PluginModelFactory
31from msmodelslim.utils.config import msmodelslim_config32from 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_repo52 custom_practice_dir = msmodelslim_config.env_vars.custom_practice_repo
52 custom_practice_path = Path(custom_practice_dir) if custom_practice_dir else None53 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_path57+ 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"""
21from pathlib import Path21from pathlib import Path
22-from typing import Dict, Generator, Optional22+from typing import Dict, Generator, List, Optional
23 23 
24from msmodelslim.app.auto_tuning import PracticeManagerInfra as atpm24from msmodelslim.app.auto_tuning import PracticeManagerInfra as atpm
25from msmodelslim.app.naive_quantization import PracticeManagerInfra as nqpm25from 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_dir43 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_databases80+ 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_config105 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))