Skip to content

Commit bdfc4c0

Browse files
committed
perf: lazy loading for TensorFlow trainers - CLI now instant
1 parent b53adf4 commit bdfc4c0

3 files changed

Lines changed: 205 additions & 39 deletions

File tree

mlcli/__init__.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -6,13 +6,23 @@
66
"""
77

88

9-
__version__="0.1.0"
10-
__author__="Devarshi Lalani"
11-
__licence__="MIT"
9+
__version__ = "0.1.0"
10+
__author__ = "Devarshi Lalani"
11+
__licence__ = "MIT"
1212

1313
from mlcli.utils.registry import ModelRegistry
1414

1515
# Global model registry instance
16+
registry = ModelRegistry()
1617

17-
registry=ModelRegistry()
18-
__all__=["registry","__version__"]
18+
19+
def _register_models():
20+
"""Register all models lazily without importing heavy dependencies."""
21+
from mlcli.trainers import register_all_models
22+
register_all_models()
23+
24+
25+
# Register models on first access
26+
_register_models()
27+
28+
__all__ = ["registry", "__version__"]

mlcli/trainers/__init__.py

Lines changed: 105 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,108 @@
1-
"""Model trainers module."""
1+
"""Model trainers module with lazy loading to avoid slow imports."""
22

3-
from mlcli.trainers.base_trainer import BaseTrainer
3+
from mlcli.trainers.base_trainer import BaseTrainer
4+
5+
# Lazy imports - TensorFlow trainers are only loaded when accessed
6+
_LAZY_IMPORTS = {
7+
"LogisticRegressionTrainer": "mlcli.trainers.logistic_trainer",
8+
"SVMTrainer": "mlcli.trainers.svm_trainer",
9+
"RFTrainer": "mlcli.trainers.rf_trainer",
10+
"XGBTrainer": "mlcli.trainers.xgb_trainer",
11+
"TFDNNTrainer": "mlcli.trainers.tf_dnn_trainer",
12+
"TFCNNTrainer": "mlcli.trainers.tf_cnn_trainer",
13+
"TFRNNTrainer": "mlcli.trainers.tf_rnn_trainer",
14+
}
15+
16+
# Pre-register models without importing heavy dependencies
17+
_MODEL_METADATA = {
18+
"logistic_regression": {
19+
"class": "LogisticRegressionTrainer",
20+
"module": "mlcli.trainers.logistic_trainer",
21+
"description": "Logistic Regression classifier with L2 regularization",
22+
"framework": "sklearn",
23+
"model_type": "classification",
24+
},
25+
"svm": {
26+
"class": "SVMTrainer",
27+
"module": "mlcli.trainers.svm_trainer",
28+
"description": "Support Vector Machine with RBF/Linear/Poly kernels",
29+
"framework": "sklearn",
30+
"model_type": "classification",
31+
},
32+
"random_forest": {
33+
"class": "RFTrainer",
34+
"module": "mlcli.trainers.rf_trainer",
35+
"description": "Random Forest ensemble classifier",
36+
"framework": "sklearn",
37+
"model_type": "classification",
38+
},
39+
"xgboost": {
40+
"class": "XGBTrainer",
41+
"module": "mlcli.trainers.xgb_trainer",
42+
"description": "XGBoost gradient boosting classifier",
43+
"framework": "xgboost",
44+
"model_type": "classification",
45+
},
46+
"tf_dnn": {
47+
"class": "TFDNNTrainer",
48+
"module": "mlcli.trainers.tf_dnn_trainer",
49+
"description": "Tensorflow Dense Feedforward Neural Network",
50+
"framework": "tensorflow",
51+
"model_type": "classification",
52+
},
53+
"tf_cnn": {
54+
"class": "TFCNNTrainer",
55+
"module": "mlcli.trainers.tf_cnn_trainer",
56+
"description": "TensorFlow Convolutional Neural Network for image classification",
57+
"framework": "tensorflow",
58+
"model_type": "classification",
59+
},
60+
"tf_rnn": {
61+
"class": "TFRNNTrainer",
62+
"module": "mlcli.trainers.tf_rnn_trainer",
63+
"description": "TensorFlow RNN/LSTM/GRU for sequence classification",
64+
"framework": "tensorflow",
65+
"model_type": "classification",
66+
},
67+
}
68+
69+
70+
def __getattr__(name: str):
71+
"""Lazy import trainers only when accessed."""
72+
if name in _LAZY_IMPORTS:
73+
import importlib
74+
module = importlib.import_module(_LAZY_IMPORTS[name])
75+
return getattr(module, name)
76+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
77+
78+
79+
def register_all_models():
80+
"""Register all models in the registry without importing heavy modules."""
81+
from mlcli import registry
82+
83+
for model_name, meta in _MODEL_METADATA.items():
84+
if not registry.is_registered(model_name):
85+
# Register with lazy loader
86+
registry.register_lazy(
87+
name=model_name,
88+
module_path=meta["module"],
89+
class_name=meta["class"],
90+
description=meta["description"],
91+
framework=meta["framework"],
92+
model_type=meta["model_type"],
93+
)
94+
95+
96+
def get_trainer_class(model_type: str):
97+
"""Get trainer class by model type, importing only when needed."""
98+
if model_type not in _MODEL_METADATA:
99+
raise ValueError(f"Unknown model type: {model_type}")
100+
101+
import importlib
102+
meta = _MODEL_METADATA[model_type]
103+
module = importlib.import_module(meta["module"])
104+
return getattr(module, meta["class"])
4105

5-
# Import all trainers to trigger auto-registration
6-
from mlcli.trainers.logistic_trainer import LogisticRegressionTrainer
7-
from mlcli.trainers.svm_trainer import SVMTrainer
8-
from mlcli.trainers.rf_trainer import RFTrainer
9-
from mlcli.trainers.xgb_trainer import XGBTrainer
10-
from mlcli.trainers.tf_dnn_trainer import TFDNNTrainer
11-
from mlcli.trainers.tf_cnn_trainer import TFCNNTrainer
12-
from mlcli.trainers.tf_rnn_trainer import TFRNNTrainer
13106

14107
__all__ = [
15108
"BaseTrainer",
@@ -20,4 +113,6 @@
20113
"TFDNNTrainer",
21114
"TFCNNTrainer",
22115
"TFRNNTrainer",
116+
"register_all_models",
117+
"get_trainer_class",
23118
]

mlcli/utils/registry.py

Lines changed: 85 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
44
Provides decorator-based auto-registration for all trainer classes,
55
enabling dynamic model discovery and instantiation from configuration.
6+
Supports lazy loading to avoid importing heavy dependencies like TensorFlow.
67
"""
78

89
from 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

Comments
 (0)