mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
Add centralized MLModelRegistry to share ML model instances
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
This commit is contained in:
parent
67d7db87cc
commit
f20148081c
8 changed files with 422 additions and 51 deletions
10
headroom/cache/dynamic_detector.py
vendored
10
headroom/cache/dynamic_detector.py
vendored
|
|
@ -595,7 +595,10 @@ class NERDetector:
|
|||
return
|
||||
|
||||
try:
|
||||
self._nlp = spacy.load(config.spacy_model)
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
self._nlp = MLModelRegistry.get_spacy(config.spacy_model)
|
||||
except OSError:
|
||||
self._load_error = (
|
||||
f"spaCy model '{config.spacy_model}' not found. "
|
||||
|
|
@ -717,7 +720,10 @@ class SemanticDetector:
|
|||
return
|
||||
|
||||
try:
|
||||
self._model = SentenceTransformer(config.embedding_model)
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
self._model = MLModelRegistry.get_sentence_transformer(config.embedding_model)
|
||||
# Pre-compute exemplar embeddings
|
||||
self._exemplar_embeddings = self._model.encode(
|
||||
self.DYNAMIC_EXEMPLARS,
|
||||
|
|
|
|||
|
|
@ -159,14 +159,16 @@ def compute_semantic_similarity(
|
|||
"""
|
||||
try:
|
||||
import numpy as np
|
||||
from sentence_transformers import SentenceTransformer
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"sentence-transformers required for semantic similarity. "
|
||||
"Install with: pip install sentence-transformers"
|
||||
) from e
|
||||
|
||||
model = SentenceTransformer(model_name)
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
model = MLModelRegistry.get_sentence_transformer(model_name)
|
||||
|
||||
embeddings = model.encode([response_a, response_b])
|
||||
embedding_a, embedding_b = embeddings[0], embeddings[1]
|
||||
|
|
|
|||
|
|
@ -18,12 +18,6 @@ from typing import Any
|
|||
|
||||
import torch
|
||||
from PIL import Image
|
||||
from transformers import (
|
||||
AutoModel,
|
||||
AutoModelForSequenceClassification,
|
||||
AutoProcessor,
|
||||
AutoTokenizer,
|
||||
)
|
||||
|
||||
|
||||
class Technique(Enum):
|
||||
|
|
@ -146,17 +140,22 @@ class TrainedRouter:
|
|||
else:
|
||||
model_id = self.DEFAULT_HF_MODEL
|
||||
|
||||
# Load classifier
|
||||
self._tokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||
self._classifier = AutoModelForSequenceClassification.from_pretrained(model_id)
|
||||
self._classifier.to(self.device) # type: ignore[attr-defined]
|
||||
self._classifier.eval() # type: ignore[attr-defined]
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
self._classifier, self._tokenizer = MLModelRegistry.get_technique_router(
|
||||
model_path=model_id,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
if self.use_siglip and self._siglip_model is None:
|
||||
self._siglip_model = AutoModel.from_pretrained(self.SIGLIP_MODEL)
|
||||
self._siglip_processor = AutoProcessor.from_pretrained(self.SIGLIP_MODEL)
|
||||
self._siglip_model.to(self.device) # type: ignore[attr-defined]
|
||||
self._siglip_model.eval() # type: ignore[attr-defined]
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
self._siglip_model, self._siglip_processor = MLModelRegistry.get_siglip(
|
||||
model_name=self.SIGLIP_MODEL,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
# Pre-compute text embeddings for image analysis
|
||||
self._compute_text_embeddings()
|
||||
|
|
|
|||
|
|
@ -137,12 +137,12 @@ class LocalEmbedder:
|
|||
return "cpu"
|
||||
|
||||
def _load_model(self) -> None:
|
||||
"""Load the sentence-transformers model lazily."""
|
||||
"""Load the sentence-transformers model lazily via MLModelRegistry."""
|
||||
if self._model is not None:
|
||||
return
|
||||
|
||||
self._check_dependencies()
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
# Determine device
|
||||
if self._requested_device:
|
||||
|
|
@ -150,13 +150,13 @@ class LocalEmbedder:
|
|||
else:
|
||||
self._device = self._detect_device()
|
||||
|
||||
logger.info(f"Loading model {self._model_name} on device {self._device}")
|
||||
self._model = SentenceTransformer(self._model_name, device=self._device)
|
||||
# Use centralized registry for shared model instances
|
||||
self._model = MLModelRegistry.get_sentence_transformer(self._model_name, self._device)
|
||||
|
||||
# Get actual dimension from loaded model
|
||||
self._dimension = self._model.get_sentence_embedding_dimension()
|
||||
logger.info(
|
||||
f"Model loaded: {self._model_name}, dimension={self._dimension}, device={self._device}"
|
||||
f"Model loaded (shared): {self._model_name}, dimension={self._dimension}, device={self._device}"
|
||||
)
|
||||
|
||||
async def embed(self, text: str) -> np.ndarray:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,9 @@
|
|||
Provides a centralized registry of LLM models with their capabilities,
|
||||
context limits, pricing, and provider information.
|
||||
|
||||
Also provides MLModelRegistry for sharing ML model instances (sentence
|
||||
transformers, SIGLIP, spaCy) to avoid loading the same model multiple times.
|
||||
|
||||
Usage:
|
||||
from headroom.models import ModelRegistry, get_model_info
|
||||
|
||||
|
|
@ -20,8 +23,18 @@ Usage:
|
|||
provider="custom",
|
||||
context_window=32000,
|
||||
)
|
||||
|
||||
# Get shared ML model instances
|
||||
from headroom.models import MLModelRegistry
|
||||
model = MLModelRegistry.get_sentence_transformer()
|
||||
"""
|
||||
|
||||
from .ml_models import (
|
||||
MLModelRegistry,
|
||||
get_sentence_transformer,
|
||||
get_siglip,
|
||||
get_spacy,
|
||||
)
|
||||
from .registry import (
|
||||
ModelInfo,
|
||||
ModelRegistry,
|
||||
|
|
@ -31,9 +44,15 @@ from .registry import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
# LLM Registry
|
||||
"ModelRegistry",
|
||||
"ModelInfo",
|
||||
"get_model_info",
|
||||
"list_models",
|
||||
"register_model",
|
||||
# ML Model Registry
|
||||
"MLModelRegistry",
|
||||
"get_sentence_transformer",
|
||||
"get_siglip",
|
||||
"get_spacy",
|
||||
]
|
||||
|
|
|
|||
361
headroom/models/ml_models.py
Normal file
361
headroom/models/ml_models.py
Normal file
|
|
@ -0,0 +1,361 @@
|
|||
"""Centralized registry for ML model instances.
|
||||
|
||||
Provides shared access to ML models (sentence transformers, SIGLIP, spaCy, etc.)
|
||||
to avoid loading the same model multiple times across different components.
|
||||
|
||||
This is different from registry.py which stores LLM metadata. This module
|
||||
manages actual loaded model instances that consume memory.
|
||||
|
||||
Usage:
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
# Get shared sentence transformer (loads on first access)
|
||||
model = MLModelRegistry.get_sentence_transformer()
|
||||
embeddings = model.encode(["hello", "world"])
|
||||
|
||||
# Get SIGLIP for image embeddings
|
||||
siglip_model, processor = MLModelRegistry.get_siglip()
|
||||
|
||||
# Check what's loaded
|
||||
print(MLModelRegistry.loaded_models())
|
||||
print(f"Memory: {MLModelRegistry.estimated_memory_mb():.1f} MB")
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from threading import RLock
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
pass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Model size estimates in MB (approximate)
|
||||
MODEL_SIZES_MB = {
|
||||
"sentence_transformer:all-MiniLM-L6-v2": 90,
|
||||
"sentence_transformer:all-mpnet-base-v2": 420,
|
||||
"siglip:google/siglip-base-patch16-224": 400,
|
||||
"siglip:google/siglip-large-patch16-384": 1200,
|
||||
"llmlingua:microsoft/llmlingua-2-xlm-roberta-large-meetingbank": 1000,
|
||||
"spacy:en_core_web_sm": 40,
|
||||
"spacy:en_core_web_md": 120,
|
||||
"technique_router:chopratejas/technique-router": 100,
|
||||
}
|
||||
|
||||
|
||||
class MLModelRegistry:
|
||||
"""Singleton registry for shared ML model instances.
|
||||
|
||||
Provides lazy-loaded, shared access to ML models across all components.
|
||||
This prevents the same model from being loaded multiple times.
|
||||
|
||||
Thread-safe for concurrent access.
|
||||
"""
|
||||
|
||||
_instance: MLModelRegistry | None = None
|
||||
_lock = RLock()
|
||||
|
||||
def __new__(cls) -> MLModelRegistry:
|
||||
if cls._instance is None:
|
||||
with cls._lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
cls._instance._init()
|
||||
return cls._instance
|
||||
|
||||
def _init(self) -> None:
|
||||
"""Initialize the registry."""
|
||||
self._models: dict[str, Any] = {}
|
||||
self._model_lock = RLock()
|
||||
|
||||
@classmethod
|
||||
def get(cls) -> MLModelRegistry:
|
||||
"""Get the singleton instance."""
|
||||
return cls()
|
||||
|
||||
@classmethod
|
||||
def reset(cls) -> None:
|
||||
"""Reset the registry (for testing)."""
|
||||
with cls._lock:
|
||||
if cls._instance is not None:
|
||||
cls._instance._models.clear()
|
||||
cls._instance = None
|
||||
|
||||
# =========================================================================
|
||||
# Sentence Transformers
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def get_sentence_transformer(
|
||||
cls,
|
||||
model_name: str = "all-MiniLM-L6-v2",
|
||||
device: str | None = None,
|
||||
) -> Any:
|
||||
"""Get a shared SentenceTransformer instance.
|
||||
|
||||
Args:
|
||||
model_name: Model name (default: all-MiniLM-L6-v2).
|
||||
device: Device to use (cuda, mps, cpu). Auto-detected if None.
|
||||
|
||||
Returns:
|
||||
SentenceTransformer model instance.
|
||||
"""
|
||||
instance = cls.get()
|
||||
key = f"sentence_transformer:{model_name}"
|
||||
|
||||
with instance._model_lock:
|
||||
if key not in instance._models:
|
||||
logger.info(f"Loading SentenceTransformer: {model_name}")
|
||||
from sentence_transformers import SentenceTransformer
|
||||
|
||||
if device is None:
|
||||
device = cls._detect_device()
|
||||
|
||||
model = SentenceTransformer(model_name, device=device)
|
||||
instance._models[key] = model
|
||||
logger.info(f"Loaded SentenceTransformer: {model_name} on {device}")
|
||||
|
||||
return instance._models[key]
|
||||
|
||||
# =========================================================================
|
||||
# SIGLIP (Image Embeddings)
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def get_siglip(
|
||||
cls,
|
||||
model_name: str = "google/siglip-base-patch16-224",
|
||||
device: str | None = None,
|
||||
) -> tuple[Any, Any]:
|
||||
"""Get shared SIGLIP model and processor.
|
||||
|
||||
Args:
|
||||
model_name: Model name (default: google/siglip-base-patch16-224).
|
||||
device: Device to use. Auto-detected if None.
|
||||
|
||||
Returns:
|
||||
Tuple of (model, processor).
|
||||
"""
|
||||
instance = cls.get()
|
||||
key = f"siglip:{model_name}"
|
||||
|
||||
with instance._model_lock:
|
||||
if key not in instance._models:
|
||||
logger.info(f"Loading SIGLIP: {model_name}")
|
||||
from transformers import AutoModel, AutoProcessor
|
||||
|
||||
if device is None:
|
||||
device = cls._detect_device()
|
||||
|
||||
model = AutoModel.from_pretrained(model_name)
|
||||
processor = AutoProcessor.from_pretrained(model_name)
|
||||
|
||||
# Move to device and set eval mode
|
||||
if device != "cpu":
|
||||
import torch
|
||||
|
||||
model = model.to(torch.device(device))
|
||||
model.eval()
|
||||
|
||||
instance._models[key] = (model, processor)
|
||||
logger.info(f"Loaded SIGLIP: {model_name} on {device}")
|
||||
|
||||
result: tuple[Any, Any] = instance._models[key]
|
||||
return result
|
||||
|
||||
# =========================================================================
|
||||
# spaCy
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def get_spacy(cls, model_name: str = "en_core_web_sm") -> Any:
|
||||
"""Get a shared spaCy model.
|
||||
|
||||
Args:
|
||||
model_name: Model name (default: en_core_web_sm).
|
||||
|
||||
Returns:
|
||||
spaCy Language model.
|
||||
"""
|
||||
instance = cls.get()
|
||||
key = f"spacy:{model_name}"
|
||||
|
||||
with instance._model_lock:
|
||||
if key not in instance._models:
|
||||
logger.info(f"Loading spaCy: {model_name}")
|
||||
import spacy
|
||||
|
||||
model = spacy.load(model_name)
|
||||
instance._models[key] = model
|
||||
logger.info(f"Loaded spaCy: {model_name}")
|
||||
|
||||
return instance._models[key]
|
||||
|
||||
# =========================================================================
|
||||
# Technique Router (Sequence Classification)
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def get_technique_router(
|
||||
cls,
|
||||
model_path: str | None = None,
|
||||
device: str | None = None,
|
||||
) -> tuple[Any, Any]:
|
||||
"""Get shared technique router model and tokenizer.
|
||||
|
||||
Args:
|
||||
model_path: Path to model (default: chopratejas/technique-router).
|
||||
device: Device to use. Auto-detected if None.
|
||||
|
||||
Returns:
|
||||
Tuple of (model, tokenizer).
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
instance = cls.get()
|
||||
|
||||
# Default to HuggingFace model, check for local first
|
||||
if model_path is None:
|
||||
local_path = Path("headroom/models/technique-router-mini/final/")
|
||||
if local_path.exists():
|
||||
model_path = str(local_path)
|
||||
else:
|
||||
model_path = "chopratejas/technique-router"
|
||||
|
||||
key = f"technique_router:{model_path}"
|
||||
|
||||
with instance._model_lock:
|
||||
if key not in instance._models:
|
||||
logger.info(f"Loading technique router: {model_path}")
|
||||
from transformers import AutoModelForSequenceClassification, AutoTokenizer
|
||||
|
||||
if device is None:
|
||||
device = cls._detect_device()
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
model = AutoModelForSequenceClassification.from_pretrained(model_path)
|
||||
|
||||
# Move to device and set eval mode
|
||||
if device != "cpu":
|
||||
import torch
|
||||
|
||||
model = model.to(torch.device(device))
|
||||
model.eval()
|
||||
|
||||
instance._models[key] = (model, tokenizer)
|
||||
logger.info(f"Loaded technique router: {model_path} on {device}")
|
||||
|
||||
result: tuple[Any, Any] = instance._models[key]
|
||||
return result
|
||||
|
||||
# =========================================================================
|
||||
# LLMLingua (uses existing singleton pattern)
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def get_llmlingua(cls, device: str | None = None, model_name: str | None = None) -> Any:
|
||||
"""Get the LLMLingua compressor.
|
||||
|
||||
Note: LLMLingua already has its own singleton in llmlingua_compressor.py.
|
||||
This method delegates to that implementation.
|
||||
|
||||
Args:
|
||||
device: Device to use. Auto-detected if None.
|
||||
model_name: Model name (default: microsoft/llmlingua-2-xlm-roberta-large-meetingbank).
|
||||
|
||||
Returns:
|
||||
PromptCompressor instance.
|
||||
"""
|
||||
from headroom.transforms.llmlingua_compressor import _get_llmlingua_compressor
|
||||
|
||||
if device is None:
|
||||
device = cls._detect_device()
|
||||
|
||||
if model_name is None:
|
||||
model_name = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank"
|
||||
|
||||
return _get_llmlingua_compressor(model_name=model_name, device=device)
|
||||
|
||||
# =========================================================================
|
||||
# Utility Methods
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def _detect_device(cls) -> str:
|
||||
"""Auto-detect the best available device."""
|
||||
try:
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
return "cuda"
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
return "mps"
|
||||
except ImportError:
|
||||
pass
|
||||
return "cpu"
|
||||
|
||||
@classmethod
|
||||
def loaded_models(cls) -> list[str]:
|
||||
"""Get list of currently loaded model keys."""
|
||||
instance = cls.get()
|
||||
with instance._model_lock:
|
||||
return list(instance._models.keys())
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls, key: str) -> bool:
|
||||
"""Check if a model is loaded."""
|
||||
instance = cls.get()
|
||||
with instance._model_lock:
|
||||
return key in instance._models
|
||||
|
||||
@classmethod
|
||||
def estimated_memory_mb(cls) -> float:
|
||||
"""Estimate total memory used by loaded models."""
|
||||
instance = cls.get()
|
||||
total = 0.0
|
||||
with instance._model_lock:
|
||||
for key in instance._models:
|
||||
total += MODEL_SIZES_MB.get(key, 100) # Default 100MB if unknown
|
||||
return total
|
||||
|
||||
@classmethod
|
||||
def get_memory_stats(cls) -> dict[str, Any]:
|
||||
"""Get memory statistics for all loaded models."""
|
||||
instance = cls.get()
|
||||
loaded_models: list[dict[str, Any]] = []
|
||||
total_estimated_mb: float = 0.0
|
||||
|
||||
with instance._model_lock:
|
||||
for key in instance._models:
|
||||
size_mb = MODEL_SIZES_MB.get(key, 100)
|
||||
loaded_models.append({"key": key, "size_mb": size_mb})
|
||||
total_estimated_mb += size_mb
|
||||
|
||||
return {
|
||||
"loaded_models": loaded_models,
|
||||
"total_estimated_mb": total_estimated_mb,
|
||||
}
|
||||
|
||||
|
||||
# Convenience functions for direct access
|
||||
def get_sentence_transformer(
|
||||
model_name: str = "all-MiniLM-L6-v2",
|
||||
device: str | None = None,
|
||||
) -> Any:
|
||||
"""Get a shared SentenceTransformer instance."""
|
||||
return MLModelRegistry.get_sentence_transformer(model_name, device)
|
||||
|
||||
|
||||
def get_siglip(
|
||||
model_name: str = "google/siglip-base-patch16-224",
|
||||
device: str | None = None,
|
||||
) -> tuple[Any, Any]:
|
||||
"""Get shared SIGLIP model and processor."""
|
||||
return MLModelRegistry.get_siglip(model_name, device)
|
||||
|
||||
|
||||
def get_spacy(model_name: str = "en_core_web_sm") -> Any:
|
||||
"""Get a shared spaCy model."""
|
||||
return MLModelRegistry.get_spacy(model_name)
|
||||
|
|
@ -1741,9 +1741,10 @@ class SemanticExtractor(BaseFeatureExtractor):
|
|||
"""Extract named entities using spaCy."""
|
||||
try:
|
||||
if self._nlp is None:
|
||||
import spacy
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
self._nlp = spacy.load("en_core_web_sm")
|
||||
self._nlp = MLModelRegistry.get_spacy("en_core_web_sm")
|
||||
|
||||
assert self._nlp is not None
|
||||
doc = self._nlp(text)
|
||||
|
|
@ -1988,10 +1989,10 @@ class EmbeddingExtractor(BaseFeatureExtractor):
|
|||
"Install with: pip install sentence-transformers"
|
||||
)
|
||||
|
||||
from sentence_transformers import SentenceTransformer
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
logger.info(f"Loading sentence transformer: {self.model_name}")
|
||||
self._model = SentenceTransformer(self.model_name, device=self.device)
|
||||
self._model = MLModelRegistry.get_sentence_transformer(self.model_name, self.device)
|
||||
return self._model
|
||||
|
||||
def extract(self, text: str, **kwargs: Any) -> EmbeddingFeatures:
|
||||
|
|
|
|||
|
|
@ -90,13 +90,11 @@ class EmbeddingScorer(RelevanceScorer):
|
|||
Requires sentence-transformers: pip install headroom[relevance]
|
||||
"""
|
||||
|
||||
_model_cache: dict[str, SentenceTransformer] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "all-MiniLM-L6-v2",
|
||||
device: str | None = None,
|
||||
cache_model: bool = True,
|
||||
cache_model: bool = True, # Kept for API compatibility, always uses registry now
|
||||
):
|
||||
"""Initialize embedding scorer.
|
||||
|
||||
|
|
@ -107,12 +105,11 @@ class EmbeddingScorer(RelevanceScorer):
|
|||
- "all-mpnet-base-v2": Best quality, slower
|
||||
- "paraphrase-MiniLM-L6-v2": Good for paraphrase detection
|
||||
device: Device to use ('cpu', 'cuda', 'mps', or None for auto).
|
||||
cache_model: If True, cache loaded models across instances.
|
||||
cache_model: Deprecated, models are always cached via MLModelRegistry.
|
||||
"""
|
||||
self.model_name = model_name
|
||||
self.device = device
|
||||
self.cache_model = cache_model
|
||||
self._model: SentenceTransformer | None = None
|
||||
self._available: bool | None = None
|
||||
|
||||
@classmethod
|
||||
|
|
@ -138,30 +135,16 @@ class EmbeddingScorer(RelevanceScorer):
|
|||
Raises:
|
||||
RuntimeError: If sentence-transformers is not installed.
|
||||
"""
|
||||
if self._model is not None:
|
||||
return self._model
|
||||
|
||||
if not self.is_available():
|
||||
raise RuntimeError(
|
||||
"EmbeddingScorer requires sentence-transformers. "
|
||||
"Install with: pip install headroom[relevance]"
|
||||
)
|
||||
|
||||
# Check cache
|
||||
if self.cache_model and self.model_name in self._model_cache:
|
||||
self._model = self._model_cache[self.model_name]
|
||||
return self._model
|
||||
# Use centralized registry for shared model instances
|
||||
from headroom.models.ml_models import MLModelRegistry
|
||||
|
||||
# Load model
|
||||
from sentence_transformers import SentenceTransformer
|
||||
|
||||
logger.info(f"Loading sentence transformer model: {self.model_name}")
|
||||
self._model = SentenceTransformer(self.model_name, device=self.device)
|
||||
|
||||
if self.cache_model:
|
||||
self._model_cache[self.model_name] = self._model
|
||||
|
||||
return self._model
|
||||
return MLModelRegistry.get_sentence_transformer(self.model_name, self.device)
|
||||
|
||||
def _encode(self, texts: list[str]):
|
||||
"""Encode texts to embeddings.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue