33
44Provides decorator-based auto-registration for all trainer classes,
55enabling dynamic model discovery and instantiation from configuration.
6+ Supports lazy loading to avoid importing heavy dependencies like TensorFlow.
67"""
78
89from typing import Dict , Type , Optional , List , Any
@@ -16,15 +17,18 @@ class ModelRegistry:
1617 Central registry for all model trainers.
1718
1819 Maps model type strings (e.g., 'logistic_regression') to their corresponding
19- trainer class implementations. Supports automatic registration via decorator.
20+ trainer class implementations. Supports automatic registration via decorator
21+ and lazy loading for heavy dependencies.
2022 """
2123
2224 def __init__ (self ) -> None :
2325 """Initialize empty registry."""
24- self ._registry :Dict [str ,Type ]= {}
25- self ._metadata : Dict [str ,Dict [str ,Any ]]= {}
26+ self ._registry : Dict [str , Type ] = {}
27+ self ._lazy_registry : Dict [str , Dict [str , str ]] = {} # For lazy loading
28+ self ._metadata : Dict [str , Dict [str , Any ]] = {}
2629
27- def register (self ,name :str ,trainer_class :Type ,description :str = "" ,framework :str = "unknown" ,model_type :str = "unknown" )-> None :
30+ def register (self , name : str , trainer_class : Type , description : str = "" ,
31+ framework : str = "unknown" , model_type : str = "unknown" ) -> None :
2832 """
2933 Register a trainer class with metadata.
3034
@@ -42,17 +46,64 @@ def register(self,name:str,trainer_class:Type,description:str="",framework:str="
4246 if name in self ._registry :
4347 logger .warning (f"Model '{ name } ' is already registered. Overwriting" )
4448
45- self ._registry [name ]= trainer_class
46- self ._metadata [name ]= {
47- "description" :description ,
48- "framework" :framework ,
49- "model_type" :model_type ,
50- "class_name" :trainer_class .__name__
49+ self ._registry [name ] = trainer_class
50+ self ._metadata [name ] = {
51+ "description" : description ,
52+ "framework" : framework ,
53+ "model_type" : model_type ,
54+ "class_name" : trainer_class .__name__
5155 }
5256
53- logger .debug (f"Registered model:{ name } -> { trainer_class .__name__ } " )
57+ logger .debug (f"Registered model: { name } -> { trainer_class .__name__ } " )
5458
55- def get (self ,name :str )-> Optional [Type ]:
59+ def register_lazy (self , name : str , module_path : str , class_name : str ,
60+ description : str = "" , framework : str = "unknown" ,
61+ model_type : str = "unknown" ) -> None :
62+ """
63+ Register a trainer for lazy loading (doesn't import the module yet).
64+
65+ Args:
66+ name: Unique identifier for the model
67+ module_path: Full module path (e.g., 'mlcli.trainers.tf_dnn_trainer')
68+ class_name: Class name to import from module
69+ description: Human-readable description
70+ framework: ML framework name
71+ model_type: Type of model
72+ """
73+ self ._lazy_registry [name ] = {
74+ "module_path" : module_path ,
75+ "class_name" : class_name ,
76+ }
77+ self ._metadata [name ] = {
78+ "description" : description ,
79+ "framework" : framework ,
80+ "model_type" : model_type ,
81+ "class_name" : class_name ,
82+ "lazy" : True ,
83+ }
84+ logger .debug (f"Registered lazy model: { name } -> { module_path } .{ class_name } " )
85+
86+ def _resolve_lazy (self , name : str ) -> Optional [Type ]:
87+ """Resolve a lazy-registered model by importing it."""
88+ if name not in self ._lazy_registry :
89+ return None
90+
91+ import importlib
92+ lazy_info = self ._lazy_registry [name ]
93+ module = importlib .import_module (lazy_info ["module_path" ])
94+ trainer_class = getattr (module , lazy_info ["class_name" ])
95+
96+ # Move from lazy to regular registry
97+ self ._registry [name ] = trainer_class
98+ del self ._lazy_registry [name ]
99+
100+ # Update metadata
101+ if name in self ._metadata :
102+ self ._metadata [name ]["lazy" ] = False
103+
104+ return trainer_class
105+
106+ def get (self , name : str ) -> Optional [Type ]:
56107 """
57108 Retrieve a trainer class by name.
58109
@@ -62,9 +113,17 @@ def get(self,name:str)->Optional[Type]:
62113 Returns:
63114 Trainer class or None if not found
64115 """
65- return self ._registry .get (name )
66-
67- def get_trainer (self ,name :str ,** kwargs )-> Any :
116+ # Check regular registry first
117+ if name in self ._registry :
118+ return self ._registry .get (name )
119+
120+ # Try lazy loading
121+ if name in self ._lazy_registry :
122+ return self ._resolve_lazy (name )
123+
124+ return None
125+
126+ def get_trainer (self , name : str , ** kwargs ) -> Any :
68127 """
69128 Instantiate a trainer by name.
70129
@@ -78,21 +137,23 @@ def get_trainer(self,name:str,**kwargs)->Any:
78137 Raises:
79138 KeyError: If model name not found in registry
80139 """
81- trainer_class = self .get (name )
140+ trainer_class = self .get (name )
82141 if trainer_class is None :
83- available = ", " .join (self .list_models ())
84- raise KeyError (f"Model '{ name } ' not found in registry." f"Available models: { available } " )
142+ available = ", " .join (self .list_models ())
143+ raise KeyError (f"Model '{ name } ' not found in registry. "
144+ f"Available models: { available } " )
85145
86146 return trainer_class (** kwargs )
87147
88- def list_models (self )-> List [str ]:
148+ def list_models (self ) -> List [str ]:
89149 """
90- Get list of all registered model names.
150+ Get list of all registered model names (including lazy) .
91151
92152 Returns:
93153 List of model identifiers
94154 """
95- return sorted (self ._registry .keys ())
155+ all_models = set (self ._registry .keys ()) | set (self ._lazy_registry .keys ())
156+ return sorted (all_models )
96157
97158
98159 def get_metadata (self ,name :str )-> Optional [Dict [str ,Any ]]:
@@ -127,17 +188,17 @@ def get_models_by_framework(self,framework:str)->List[str]:
127188 """
128189 return [name for name ,meta in self ._metadata .items () if meta .get ("framework" )== framework ]
129190
130- def is_registered (self ,name :str )-> bool :
191+ def is_registered (self , name : str ) -> bool :
131192 """
132- Check if a model is registered.
193+ Check if a model is registered (including lazy) .
133194
134195 Args:
135196 name: Model identifier
136197
137198 Returns:
138199 True if registered, False otherwise
139200 """
140- return name in self ._registry
201+ return name in self ._registry or name in self . _lazy_registry
141202
142203 def unregister (self ,name :str )-> bool :
143204 """
0 commit comments