chopratejas commited on
Commit
bb32b57
·
1 Parent(s): ef87114

Add centralized MLModelRegistry to share ML model instances

Browse files

Previously, SentenceTransformer was loaded up to 5 times in different
components, wasting ~1.5GB of memory. Now all ML models are shared via
MLModelRegistry:

- SentenceTransformer (text embeddings)
- SIGLIP (image embeddings)
- spaCy (NER)
- Technique router (image optimization)

Updated components to use the registry:
- headroom/relevance/embedding.py
- headroom/memory/adapters/embedders.py
- headroom/cache/dynamic_detector.py
- headroom/prediction/feature_extractor.py
- headroom/evals/metrics.py
- headroom/image/trained_router.py

headroom/cache/dynamic_detector.py CHANGED
@@ -595,7 +595,10 @@ class NERDetector:
595
  return
596
 
597
  try:
598
- self._nlp = spacy.load(config.spacy_model)
 
 
 
599
  except OSError:
600
  self._load_error = (
601
  f"spaCy model '{config.spacy_model}' not found. "
@@ -717,7 +720,10 @@ class SemanticDetector:
717
  return
718
 
719
  try:
720
- self._model = SentenceTransformer(config.embedding_model)
 
 
 
721
  # Pre-compute exemplar embeddings
722
  self._exemplar_embeddings = self._model.encode(
723
  self.DYNAMIC_EXEMPLARS,
 
595
  return
596
 
597
  try:
598
+ # Use centralized registry for shared model instances
599
+ from headroom.models.ml_models import MLModelRegistry
600
+
601
+ self._nlp = MLModelRegistry.get_spacy(config.spacy_model)
602
  except OSError:
603
  self._load_error = (
604
  f"spaCy model '{config.spacy_model}' not found. "
 
720
  return
721
 
722
  try:
723
+ # Use centralized registry for shared model instances
724
+ from headroom.models.ml_models import MLModelRegistry
725
+
726
+ self._model = MLModelRegistry.get_sentence_transformer(config.embedding_model)
727
  # Pre-compute exemplar embeddings
728
  self._exemplar_embeddings = self._model.encode(
729
  self.DYNAMIC_EXEMPLARS,
headroom/evals/metrics.py CHANGED
@@ -159,14 +159,16 @@ def compute_semantic_similarity(
159
  """
160
  try:
161
  import numpy as np
162
- from sentence_transformers import SentenceTransformer
163
  except ImportError as e:
164
  raise ImportError(
165
  "sentence-transformers required for semantic similarity. "
166
  "Install with: pip install sentence-transformers"
167
  ) from e
168
 
169
- model = SentenceTransformer(model_name)
 
 
 
170
 
171
  embeddings = model.encode([response_a, response_b])
172
  embedding_a, embedding_b = embeddings[0], embeddings[1]
 
159
  """
160
  try:
161
  import numpy as np
 
162
  except ImportError as e:
163
  raise ImportError(
164
  "sentence-transformers required for semantic similarity. "
165
  "Install with: pip install sentence-transformers"
166
  ) from e
167
 
168
+ # Use centralized registry for shared model instances
169
+ from headroom.models.ml_models import MLModelRegistry
170
+
171
+ model = MLModelRegistry.get_sentence_transformer(model_name)
172
 
173
  embeddings = model.encode([response_a, response_b])
174
  embedding_a, embedding_b = embeddings[0], embeddings[1]
headroom/image/trained_router.py CHANGED
@@ -18,12 +18,6 @@ from typing import Any
18
 
19
  import torch
20
  from PIL import Image
21
- from transformers import (
22
- AutoModel,
23
- AutoModelForSequenceClassification,
24
- AutoProcessor,
25
- AutoTokenizer,
26
- )
27
 
28
 
29
  class Technique(Enum):
@@ -146,17 +140,22 @@ class TrainedRouter:
146
  else:
147
  model_id = self.DEFAULT_HF_MODEL
148
 
149
- # Load classifier
150
- self._tokenizer = AutoTokenizer.from_pretrained(model_id)
151
- self._classifier = AutoModelForSequenceClassification.from_pretrained(model_id)
152
- self._classifier.to(self.device) # type: ignore[attr-defined]
153
- self._classifier.eval() # type: ignore[attr-defined]
 
 
154
 
155
  if self.use_siglip and self._siglip_model is None:
156
- self._siglip_model = AutoModel.from_pretrained(self.SIGLIP_MODEL)
157
- self._siglip_processor = AutoProcessor.from_pretrained(self.SIGLIP_MODEL)
158
- self._siglip_model.to(self.device) # type: ignore[attr-defined]
159
- self._siglip_model.eval() # type: ignore[attr-defined]
 
 
 
160
 
161
  # Pre-compute text embeddings for image analysis
162
  self._compute_text_embeddings()
 
18
 
19
  import torch
20
  from PIL import Image
 
 
 
 
 
 
21
 
22
 
23
  class Technique(Enum):
 
140
  else:
141
  model_id = self.DEFAULT_HF_MODEL
142
 
143
+ # Use centralized registry for shared model instances
144
+ from headroom.models.ml_models import MLModelRegistry
145
+
146
+ self._classifier, self._tokenizer = MLModelRegistry.get_technique_router(
147
+ model_path=model_id,
148
+ device=self.device,
149
+ )
150
 
151
  if self.use_siglip and self._siglip_model is None:
152
+ # Use centralized registry for shared model instances
153
+ from headroom.models.ml_models import MLModelRegistry
154
+
155
+ self._siglip_model, self._siglip_processor = MLModelRegistry.get_siglip(
156
+ model_name=self.SIGLIP_MODEL,
157
+ device=self.device,
158
+ )
159
 
160
  # Pre-compute text embeddings for image analysis
161
  self._compute_text_embeddings()
headroom/memory/adapters/embedders.py CHANGED
@@ -137,12 +137,12 @@ class LocalEmbedder:
137
  return "cpu"
138
 
139
  def _load_model(self) -> None:
140
- """Load the sentence-transformers model lazily."""
141
  if self._model is not None:
142
  return
143
 
144
  self._check_dependencies()
145
- from sentence_transformers import SentenceTransformer
146
 
147
  # Determine device
148
  if self._requested_device:
@@ -150,13 +150,13 @@ class LocalEmbedder:
150
  else:
151
  self._device = self._detect_device()
152
 
153
- logger.info(f"Loading model {self._model_name} on device {self._device}")
154
- self._model = SentenceTransformer(self._model_name, device=self._device)
155
 
156
  # Get actual dimension from loaded model
157
  self._dimension = self._model.get_sentence_embedding_dimension()
158
  logger.info(
159
- f"Model loaded: {self._model_name}, dimension={self._dimension}, device={self._device}"
160
  )
161
 
162
  async def embed(self, text: str) -> np.ndarray:
 
137
  return "cpu"
138
 
139
  def _load_model(self) -> None:
140
+ """Load the sentence-transformers model lazily via MLModelRegistry."""
141
  if self._model is not None:
142
  return
143
 
144
  self._check_dependencies()
145
+ from headroom.models.ml_models import MLModelRegistry
146
 
147
  # Determine device
148
  if self._requested_device:
 
150
  else:
151
  self._device = self._detect_device()
152
 
153
+ # Use centralized registry for shared model instances
154
+ self._model = MLModelRegistry.get_sentence_transformer(self._model_name, self._device)
155
 
156
  # Get actual dimension from loaded model
157
  self._dimension = self._model.get_sentence_embedding_dimension()
158
  logger.info(
159
+ f"Model loaded (shared): {self._model_name}, dimension={self._dimension}, device={self._device}"
160
  )
161
 
162
  async def embed(self, text: str) -> np.ndarray:
headroom/models/__init__.py CHANGED
@@ -3,6 +3,9 @@
3
  Provides a centralized registry of LLM models with their capabilities,
4
  context limits, pricing, and provider information.
5
 
 
 
 
6
  Usage:
7
  from headroom.models import ModelRegistry, get_model_info
8
 
@@ -20,8 +23,18 @@ Usage:
20
  provider="custom",
21
  context_window=32000,
22
  )
 
 
 
 
23
  """
24
 
 
 
 
 
 
 
25
  from .registry import (
26
  ModelInfo,
27
  ModelRegistry,
@@ -31,9 +44,15 @@ from .registry import (
31
  )
32
 
33
  __all__ = [
 
34
  "ModelRegistry",
35
  "ModelInfo",
36
  "get_model_info",
37
  "list_models",
38
  "register_model",
 
 
 
 
 
39
  ]
 
3
  Provides a centralized registry of LLM models with their capabilities,
4
  context limits, pricing, and provider information.
5
 
6
+ Also provides MLModelRegistry for sharing ML model instances (sentence
7
+ transformers, SIGLIP, spaCy) to avoid loading the same model multiple times.
8
+
9
  Usage:
10
  from headroom.models import ModelRegistry, get_model_info
11
 
 
23
  provider="custom",
24
  context_window=32000,
25
  )
26
+
27
+ # Get shared ML model instances
28
+ from headroom.models import MLModelRegistry
29
+ model = MLModelRegistry.get_sentence_transformer()
30
  """
31
 
32
+ from .ml_models import (
33
+ MLModelRegistry,
34
+ get_sentence_transformer,
35
+ get_siglip,
36
+ get_spacy,
37
+ )
38
  from .registry import (
39
  ModelInfo,
40
  ModelRegistry,
 
44
  )
45
 
46
  __all__ = [
47
+ # LLM Registry
48
  "ModelRegistry",
49
  "ModelInfo",
50
  "get_model_info",
51
  "list_models",
52
  "register_model",
53
+ # ML Model Registry
54
+ "MLModelRegistry",
55
+ "get_sentence_transformer",
56
+ "get_siglip",
57
+ "get_spacy",
58
  ]
headroom/models/ml_models.py ADDED
@@ -0,0 +1,361 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Centralized registry for ML model instances.
2
+
3
+ Provides shared access to ML models (sentence transformers, SIGLIP, spaCy, etc.)
4
+ to avoid loading the same model multiple times across different components.
5
+
6
+ This is different from registry.py which stores LLM metadata. This module
7
+ manages actual loaded model instances that consume memory.
8
+
9
+ Usage:
10
+ from headroom.models.ml_models import MLModelRegistry
11
+
12
+ # Get shared sentence transformer (loads on first access)
13
+ model = MLModelRegistry.get_sentence_transformer()
14
+ embeddings = model.encode(["hello", "world"])
15
+
16
+ # Get SIGLIP for image embeddings
17
+ siglip_model, processor = MLModelRegistry.get_siglip()
18
+
19
+ # Check what's loaded
20
+ print(MLModelRegistry.loaded_models())
21
+ print(f"Memory: {MLModelRegistry.estimated_memory_mb():.1f} MB")
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import logging
27
+ from threading import RLock
28
+ from typing import TYPE_CHECKING, Any
29
+
30
+ if TYPE_CHECKING:
31
+ pass
32
+
33
+ logger = logging.getLogger(__name__)
34
+
35
+ # Model size estimates in MB (approximate)
36
+ MODEL_SIZES_MB = {
37
+ "sentence_transformer:all-MiniLM-L6-v2": 90,
38
+ "sentence_transformer:all-mpnet-base-v2": 420,
39
+ "siglip:google/siglip-base-patch16-224": 400,
40
+ "siglip:google/siglip-large-patch16-384": 1200,
41
+ "llmlingua:microsoft/llmlingua-2-xlm-roberta-large-meetingbank": 1000,
42
+ "spacy:en_core_web_sm": 40,
43
+ "spacy:en_core_web_md": 120,
44
+ "technique_router:chopratejas/technique-router": 100,
45
+ }
46
+
47
+
48
+ class MLModelRegistry:
49
+ """Singleton registry for shared ML model instances.
50
+
51
+ Provides lazy-loaded, shared access to ML models across all components.
52
+ This prevents the same model from being loaded multiple times.
53
+
54
+ Thread-safe for concurrent access.
55
+ """
56
+
57
+ _instance: MLModelRegistry | None = None
58
+ _lock = RLock()
59
+
60
+ def __new__(cls) -> MLModelRegistry:
61
+ if cls._instance is None:
62
+ with cls._lock:
63
+ if cls._instance is None:
64
+ cls._instance = super().__new__(cls)
65
+ cls._instance._init()
66
+ return cls._instance
67
+
68
+ def _init(self) -> None:
69
+ """Initialize the registry."""
70
+ self._models: dict[str, Any] = {}
71
+ self._model_lock = RLock()
72
+
73
+ @classmethod
74
+ def get(cls) -> MLModelRegistry:
75
+ """Get the singleton instance."""
76
+ return cls()
77
+
78
+ @classmethod
79
+ def reset(cls) -> None:
80
+ """Reset the registry (for testing)."""
81
+ with cls._lock:
82
+ if cls._instance is not None:
83
+ cls._instance._models.clear()
84
+ cls._instance = None
85
+
86
+ # =========================================================================
87
+ # Sentence Transformers
88
+ # =========================================================================
89
+
90
+ @classmethod
91
+ def get_sentence_transformer(
92
+ cls,
93
+ model_name: str = "all-MiniLM-L6-v2",
94
+ device: str | None = None,
95
+ ) -> Any:
96
+ """Get a shared SentenceTransformer instance.
97
+
98
+ Args:
99
+ model_name: Model name (default: all-MiniLM-L6-v2).
100
+ device: Device to use (cuda, mps, cpu). Auto-detected if None.
101
+
102
+ Returns:
103
+ SentenceTransformer model instance.
104
+ """
105
+ instance = cls.get()
106
+ key = f"sentence_transformer:{model_name}"
107
+
108
+ with instance._model_lock:
109
+ if key not in instance._models:
110
+ logger.info(f"Loading SentenceTransformer: {model_name}")
111
+ from sentence_transformers import SentenceTransformer
112
+
113
+ if device is None:
114
+ device = cls._detect_device()
115
+
116
+ model = SentenceTransformer(model_name, device=device)
117
+ instance._models[key] = model
118
+ logger.info(f"Loaded SentenceTransformer: {model_name} on {device}")
119
+
120
+ return instance._models[key]
121
+
122
+ # =========================================================================
123
+ # SIGLIP (Image Embeddings)
124
+ # =========================================================================
125
+
126
+ @classmethod
127
+ def get_siglip(
128
+ cls,
129
+ model_name: str = "google/siglip-base-patch16-224",
130
+ device: str | None = None,
131
+ ) -> tuple[Any, Any]:
132
+ """Get shared SIGLIP model and processor.
133
+
134
+ Args:
135
+ model_name: Model name (default: google/siglip-base-patch16-224).
136
+ device: Device to use. Auto-detected if None.
137
+
138
+ Returns:
139
+ Tuple of (model, processor).
140
+ """
141
+ instance = cls.get()
142
+ key = f"siglip:{model_name}"
143
+
144
+ with instance._model_lock:
145
+ if key not in instance._models:
146
+ logger.info(f"Loading SIGLIP: {model_name}")
147
+ from transformers import AutoModel, AutoProcessor
148
+
149
+ if device is None:
150
+ device = cls._detect_device()
151
+
152
+ model = AutoModel.from_pretrained(model_name)
153
+ processor = AutoProcessor.from_pretrained(model_name)
154
+
155
+ # Move to device and set eval mode
156
+ if device != "cpu":
157
+ import torch
158
+
159
+ model = model.to(torch.device(device))
160
+ model.eval()
161
+
162
+ instance._models[key] = (model, processor)
163
+ logger.info(f"Loaded SIGLIP: {model_name} on {device}")
164
+
165
+ result: tuple[Any, Any] = instance._models[key]
166
+ return result
167
+
168
+ # =========================================================================
169
+ # spaCy
170
+ # =========================================================================
171
+
172
+ @classmethod
173
+ def get_spacy(cls, model_name: str = "en_core_web_sm") -> Any:
174
+ """Get a shared spaCy model.
175
+
176
+ Args:
177
+ model_name: Model name (default: en_core_web_sm).
178
+
179
+ Returns:
180
+ spaCy Language model.
181
+ """
182
+ instance = cls.get()
183
+ key = f"spacy:{model_name}"
184
+
185
+ with instance._model_lock:
186
+ if key not in instance._models:
187
+ logger.info(f"Loading spaCy: {model_name}")
188
+ import spacy
189
+
190
+ model = spacy.load(model_name)
191
+ instance._models[key] = model
192
+ logger.info(f"Loaded spaCy: {model_name}")
193
+
194
+ return instance._models[key]
195
+
196
+ # =========================================================================
197
+ # Technique Router (Sequence Classification)
198
+ # =========================================================================
199
+
200
+ @classmethod
201
+ def get_technique_router(
202
+ cls,
203
+ model_path: str | None = None,
204
+ device: str | None = None,
205
+ ) -> tuple[Any, Any]:
206
+ """Get shared technique router model and tokenizer.
207
+
208
+ Args:
209
+ model_path: Path to model (default: chopratejas/technique-router).
210
+ device: Device to use. Auto-detected if None.
211
+
212
+ Returns:
213
+ Tuple of (model, tokenizer).
214
+ """
215
+ from pathlib import Path
216
+
217
+ instance = cls.get()
218
+
219
+ # Default to HuggingFace model, check for local first
220
+ if model_path is None:
221
+ local_path = Path("headroom/models/technique-router-mini/final/")
222
+ if local_path.exists():
223
+ model_path = str(local_path)
224
+ else:
225
+ model_path = "chopratejas/technique-router"
226
+
227
+ key = f"technique_router:{model_path}"
228
+
229
+ with instance._model_lock:
230
+ if key not in instance._models:
231
+ logger.info(f"Loading technique router: {model_path}")
232
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
233
+
234
+ if device is None:
235
+ device = cls._detect_device()
236
+
237
+ tokenizer = AutoTokenizer.from_pretrained(model_path)
238
+ model = AutoModelForSequenceClassification.from_pretrained(model_path)
239
+
240
+ # Move to device and set eval mode
241
+ if device != "cpu":
242
+ import torch
243
+
244
+ model = model.to(torch.device(device))
245
+ model.eval()
246
+
247
+ instance._models[key] = (model, tokenizer)
248
+ logger.info(f"Loaded technique router: {model_path} on {device}")
249
+
250
+ result: tuple[Any, Any] = instance._models[key]
251
+ return result
252
+
253
+ # =========================================================================
254
+ # LLMLingua (uses existing singleton pattern)
255
+ # =========================================================================
256
+
257
+ @classmethod
258
+ def get_llmlingua(cls, device: str | None = None, model_name: str | None = None) -> Any:
259
+ """Get the LLMLingua compressor.
260
+
261
+ Note: LLMLingua already has its own singleton in llmlingua_compressor.py.
262
+ This method delegates to that implementation.
263
+
264
+ Args:
265
+ device: Device to use. Auto-detected if None.
266
+ model_name: Model name (default: microsoft/llmlingua-2-xlm-roberta-large-meetingbank).
267
+
268
+ Returns:
269
+ PromptCompressor instance.
270
+ """
271
+ from headroom.transforms.llmlingua_compressor import _get_llmlingua_compressor
272
+
273
+ if device is None:
274
+ device = cls._detect_device()
275
+
276
+ if model_name is None:
277
+ model_name = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank"
278
+
279
+ return _get_llmlingua_compressor(model_name=model_name, device=device)
280
+
281
+ # =========================================================================
282
+ # Utility Methods
283
+ # =========================================================================
284
+
285
+ @classmethod
286
+ def _detect_device(cls) -> str:
287
+ """Auto-detect the best available device."""
288
+ try:
289
+ import torch
290
+
291
+ if torch.cuda.is_available():
292
+ return "cuda"
293
+ if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
294
+ return "mps"
295
+ except ImportError:
296
+ pass
297
+ return "cpu"
298
+
299
+ @classmethod
300
+ def loaded_models(cls) -> list[str]:
301
+ """Get list of currently loaded model keys."""
302
+ instance = cls.get()
303
+ with instance._model_lock:
304
+ return list(instance._models.keys())
305
+
306
+ @classmethod
307
+ def is_loaded(cls, key: str) -> bool:
308
+ """Check if a model is loaded."""
309
+ instance = cls.get()
310
+ with instance._model_lock:
311
+ return key in instance._models
312
+
313
+ @classmethod
314
+ def estimated_memory_mb(cls) -> float:
315
+ """Estimate total memory used by loaded models."""
316
+ instance = cls.get()
317
+ total = 0.0
318
+ with instance._model_lock:
319
+ for key in instance._models:
320
+ total += MODEL_SIZES_MB.get(key, 100) # Default 100MB if unknown
321
+ return total
322
+
323
+ @classmethod
324
+ def get_memory_stats(cls) -> dict[str, Any]:
325
+ """Get memory statistics for all loaded models."""
326
+ instance = cls.get()
327
+ loaded_models: list[dict[str, Any]] = []
328
+ total_estimated_mb: float = 0.0
329
+
330
+ with instance._model_lock:
331
+ for key in instance._models:
332
+ size_mb = MODEL_SIZES_MB.get(key, 100)
333
+ loaded_models.append({"key": key, "size_mb": size_mb})
334
+ total_estimated_mb += size_mb
335
+
336
+ return {
337
+ "loaded_models": loaded_models,
338
+ "total_estimated_mb": total_estimated_mb,
339
+ }
340
+
341
+
342
+ # Convenience functions for direct access
343
+ def get_sentence_transformer(
344
+ model_name: str = "all-MiniLM-L6-v2",
345
+ device: str | None = None,
346
+ ) -> Any:
347
+ """Get a shared SentenceTransformer instance."""
348
+ return MLModelRegistry.get_sentence_transformer(model_name, device)
349
+
350
+
351
+ def get_siglip(
352
+ model_name: str = "google/siglip-base-patch16-224",
353
+ device: str | None = None,
354
+ ) -> tuple[Any, Any]:
355
+ """Get shared SIGLIP model and processor."""
356
+ return MLModelRegistry.get_siglip(model_name, device)
357
+
358
+
359
+ def get_spacy(model_name: str = "en_core_web_sm") -> Any:
360
+ """Get a shared spaCy model."""
361
+ return MLModelRegistry.get_spacy(model_name)
headroom/prediction/feature_extractor.py CHANGED
@@ -1741,9 +1741,10 @@ class SemanticExtractor(BaseFeatureExtractor):
1741
  """Extract named entities using spaCy."""
1742
  try:
1743
  if self._nlp is None:
1744
- import spacy
 
1745
 
1746
- self._nlp = spacy.load("en_core_web_sm")
1747
 
1748
  assert self._nlp is not None
1749
  doc = self._nlp(text)
@@ -1988,10 +1989,10 @@ class EmbeddingExtractor(BaseFeatureExtractor):
1988
  "Install with: pip install sentence-transformers"
1989
  )
1990
 
1991
- from sentence_transformers import SentenceTransformer
 
1992
 
1993
- logger.info(f"Loading sentence transformer: {self.model_name}")
1994
- self._model = SentenceTransformer(self.model_name, device=self.device)
1995
  return self._model
1996
 
1997
  def extract(self, text: str, **kwargs: Any) -> EmbeddingFeatures:
 
1741
  """Extract named entities using spaCy."""
1742
  try:
1743
  if self._nlp is None:
1744
+ # Use centralized registry for shared model instances
1745
+ from headroom.models.ml_models import MLModelRegistry
1746
 
1747
+ self._nlp = MLModelRegistry.get_spacy("en_core_web_sm")
1748
 
1749
  assert self._nlp is not None
1750
  doc = self._nlp(text)
 
1989
  "Install with: pip install sentence-transformers"
1990
  )
1991
 
1992
+ # Use centralized registry for shared model instances
1993
+ from headroom.models.ml_models import MLModelRegistry
1994
 
1995
+ self._model = MLModelRegistry.get_sentence_transformer(self.model_name, self.device)
 
1996
  return self._model
1997
 
1998
  def extract(self, text: str, **kwargs: Any) -> EmbeddingFeatures:
headroom/relevance/embedding.py CHANGED
@@ -90,13 +90,11 @@ class EmbeddingScorer(RelevanceScorer):
90
  Requires sentence-transformers: pip install headroom[relevance]
91
  """
92
 
93
- _model_cache: dict[str, SentenceTransformer] = {}
94
-
95
  def __init__(
96
  self,
97
  model_name: str = "all-MiniLM-L6-v2",
98
  device: str | None = None,
99
- cache_model: bool = True,
100
  ):
101
  """Initialize embedding scorer.
102
 
@@ -107,12 +105,11 @@ class EmbeddingScorer(RelevanceScorer):
107
  - "all-mpnet-base-v2": Best quality, slower
108
  - "paraphrase-MiniLM-L6-v2": Good for paraphrase detection
109
  device: Device to use ('cpu', 'cuda', 'mps', or None for auto).
110
- cache_model: If True, cache loaded models across instances.
111
  """
112
  self.model_name = model_name
113
  self.device = device
114
  self.cache_model = cache_model
115
- self._model: SentenceTransformer | None = None
116
  self._available: bool | None = None
117
 
118
  @classmethod
@@ -138,30 +135,16 @@ class EmbeddingScorer(RelevanceScorer):
138
  Raises:
139
  RuntimeError: If sentence-transformers is not installed.
140
  """
141
- if self._model is not None:
142
- return self._model
143
-
144
  if not self.is_available():
145
  raise RuntimeError(
146
  "EmbeddingScorer requires sentence-transformers. "
147
  "Install with: pip install headroom[relevance]"
148
  )
149
 
150
- # Check cache
151
- if self.cache_model and self.model_name in self._model_cache:
152
- self._model = self._model_cache[self.model_name]
153
- return self._model
154
-
155
- # Load model
156
- from sentence_transformers import SentenceTransformer
157
-
158
- logger.info(f"Loading sentence transformer model: {self.model_name}")
159
- self._model = SentenceTransformer(self.model_name, device=self.device)
160
-
161
- if self.cache_model:
162
- self._model_cache[self.model_name] = self._model
163
 
164
- return self._model
165
 
166
  def _encode(self, texts: list[str]):
167
  """Encode texts to embeddings.
 
90
  Requires sentence-transformers: pip install headroom[relevance]
91
  """
92
 
 
 
93
  def __init__(
94
  self,
95
  model_name: str = "all-MiniLM-L6-v2",
96
  device: str | None = None,
97
+ cache_model: bool = True, # Kept for API compatibility, always uses registry now
98
  ):
99
  """Initialize embedding scorer.
100
 
 
105
  - "all-mpnet-base-v2": Best quality, slower
106
  - "paraphrase-MiniLM-L6-v2": Good for paraphrase detection
107
  device: Device to use ('cpu', 'cuda', 'mps', or None for auto).
108
+ cache_model: Deprecated, models are always cached via MLModelRegistry.
109
  """
110
  self.model_name = model_name
111
  self.device = device
112
  self.cache_model = cache_model
 
113
  self._available: bool | None = None
114
 
115
  @classmethod
 
135
  Raises:
136
  RuntimeError: If sentence-transformers is not installed.
137
  """
 
 
 
138
  if not self.is_available():
139
  raise RuntimeError(
140
  "EmbeddingScorer requires sentence-transformers. "
141
  "Install with: pip install headroom[relevance]"
142
  )
143
 
144
+ # Use centralized registry for shared model instances
145
+ from headroom.models.ml_models import MLModelRegistry
 
 
 
 
 
 
 
 
 
 
 
146
 
147
+ return MLModelRegistry.get_sentence_transformer(self.model_name, self.device)
148
 
149
  def _encode(self, texts: list[str]):
150
  """Encode texts to embeddings.