fix: resolve model switching bug and implement intrinsic dimension detection in TextEmbedder

This commit is contained in:
KaifAhmad1
2026-01-08 18:56:42 +05:30
parent 9a2f2cd2d2
commit 2bd1d06eb2
2 changed files with 155 additions and 29 deletions
+29 -27
View File
@@ -107,6 +107,7 @@ class TextEmbedder:
# Initialize models (will be None if unavailable)
self.model = None
self.fastembed_model = None
self.embedding_dimension = None
# Initialize progress tracker
self.progress_tracker = get_progress_tracker()
@@ -121,6 +122,11 @@ class TextEmbedder:
If unavailable or loading fails, falls back to hash-based embedding method.
Logs warnings but doesn't raise errors to allow graceful degradation.
"""
# Clear previous models to avoid conflicts when switching
self.model = None
self.fastembed_model = None
self.embedding_dimension = None
if self.method == "fastembed":
if FASTEMBED_AVAILABLE:
try:
@@ -129,8 +135,17 @@ class TextEmbedder:
"fastembed_model_name", self.model_name
)
self.fastembed_model = TextEmbedding(model_name=fastembed_model_name)
# Detect dimension from model
try:
# FastEmbed doesn't expose dimension directly, so we generate a test embedding
test_emb = list(self.fastembed_model.embed(["test"]))[0]
self.embedding_dimension = len(test_emb)
except Exception:
self.embedding_dimension = 384
self.logger.info(
f"Loaded FastEmbed model: {fastembed_model_name}"
f"Loaded FastEmbed model: {fastembed_model_name} (dim: {self.embedding_dimension})"
)
except Exception as e:
self.logger.warning(
@@ -149,9 +164,10 @@ class TextEmbedder:
if SENTENCE_TRANSFORMERS_AVAILABLE:
try:
self.model = SentenceTransformer(self.model_name, device=self.device)
self.embedding_dimension = self.model.get_sentence_embedding_dimension()
self.logger.info(
f"Loaded sentence-transformers model: {self.model_name} "
f"(device: {self.device})"
f"(device: {self.device}, dim: {self.embedding_dimension})"
)
except Exception as e:
self.logger.warning(
@@ -166,18 +182,11 @@ class TextEmbedder:
"Using fallback embedding method."
)
def get_method(self) -> str:
"""Get current embedding method."""
return self.method
# If no model loaded, use fallback dimension
if self.embedding_dimension is None:
self.embedding_dimension = self.config.get("dimension", 128)
def get_model_info(self) -> Dict[str, Any]:
"""Get current model information."""
return {
"method": self.method,
"model_name": self.model_name,
"device": self.device,
"normalize": self.normalize
}
self.logger.debug(f"Initialized text embedder with dimension: {self.embedding_dimension}")
def set_model(self, method: str, model_name: str, **config) -> None:
"""
@@ -190,6 +199,10 @@ class TextEmbedder:
"""
self.method = method.lower()
self.model_name = model_name
# Update config with new values
self.config.update(config)
if "device" in config:
self.device = config["device"]
if "normalize" in config:
@@ -340,8 +353,8 @@ class TextEmbedder:
hash_val = int(hashlib.sha256(text.encode("utf-8")).hexdigest(), 16)
rng = np.random.RandomState(hash_val % (2**32))
# Determine dimension - use config or default
dim = self.config.get("dimension", 128)
# Determine dimension - use stored dimension or default
dim = self.embedding_dimension or 128
embedding = rng.rand(dim).astype(np.float32)
# Normalize if requested
@@ -470,18 +483,7 @@ class TextEmbedder:
>>> dim = embedder.get_embedding_dimension()
>>> print(f"Embedding dimension: {dim}")
"""
if self.fastembed_model:
# FastEmbed doesn't expose dimension directly, so we generate a test embedding
try:
test_emb = self._embed_with_fastembed("test")
return len(test_emb)
except Exception:
# Default FastEmbed dimension
return 384
elif self.model:
return self.model.get_sentence_embedding_dimension()
else:
return 128 # Fallback hash-based embedding dimension
return self.embedding_dimension or 128
def get_method(self) -> str:
"""
+124
View File
@@ -0,0 +1,124 @@
import unittest
from unittest.mock import MagicMock, patch
import numpy as np
import sys
import os
# Ensure the package is in the path
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '../..')))
from semantica.embeddings.text_embedder import TextEmbedder
from semantica.embeddings.embedding_generator import EmbeddingGenerator
class TestModelSwitching(unittest.TestCase):
"""Test suite for dynamic model switching in TextEmbedder and EmbeddingGenerator."""
def setUp(self):
# Mock dependencies
self.mock_st_patcher = patch('semantica.embeddings.text_embedder.SentenceTransformer')
self.mock_st_class = self.mock_st_patcher.start()
self.mock_fe_patcher = patch('semantica.embeddings.text_embedder.TextEmbedding')
self.mock_fe_class = self.mock_fe_patcher.start()
# Patch availability flags
self.st_avail_patcher = patch('semantica.embeddings.text_embedder.SENTENCE_TRANSFORMERS_AVAILABLE', True)
self.st_avail_patcher.start()
self.fe_avail_patcher = patch('semantica.embeddings.text_embedder.FASTEMBED_AVAILABLE', True)
self.fe_avail_patcher.start()
def tearDown(self):
self.mock_st_patcher.stop()
self.mock_fe_patcher.stop()
self.st_avail_patcher.stop()
self.fe_avail_patcher.stop()
def test_text_embedder_dynamic_switching(self):
"""Test that TextEmbedder correctly switches between methods and clears state."""
# 1. Start with fastembed
embedder = TextEmbedder(method="fastembed", model_name="fast-model")
self.assertEqual(embedder.get_method(), "fastembed")
self.assertIsNotNone(embedder.fastembed_model)
self.assertIsNone(embedder.model)
# 2. Switch to sentence_transformers
embedder.set_model(method="sentence_transformers", model_name="st-model")
# Verify method and model state
self.assertEqual(embedder.get_method(), "sentence_transformers")
self.assertIsNone(embedder.fastembed_model, "fastembed_model should be cleared after switching")
self.assertIsNotNone(embedder.model, "sentence_transformer model should be initialized")
self.assertEqual(embedder.model_name, "st-model")
# 3. Switch back to fastembed
embedder.set_model(method="fastembed", model_name="fast-model-v2")
self.assertEqual(embedder.get_method(), "fastembed")
self.assertIsNotNone(embedder.fastembed_model, "fastembed_model should be re-initialized")
self.assertIsNone(embedder.model, "sentence_transformer model should be cleared")
self.assertEqual(embedder.model_name, "fast-model-v2")
def test_embedding_generator_dynamic_switching(self):
"""Test that EmbeddingGenerator correctly propagates model switches to TextEmbedder."""
gen = EmbeddingGenerator()
# Default should be fastembed
self.assertEqual(gen.get_text_method(), "fastembed")
# Switch via EmbeddingGenerator
gen.set_text_model(method="sentence_transformers", model_name="new-st-model")
# Verify through Generator methods
self.assertEqual(gen.get_text_method(), "sentence_transformers")
info = gen.get_methods_info()
self.assertEqual(info['text']['method'], "sentence_transformers")
self.assertEqual(info['text']['model_name'], "new-st-model")
self.assertTrue(info['text']['model_loaded'])
def test_fallback_on_switch_failure(self):
"""Test that switching to an unavailable method falls back to 'fallback'."""
embedder = TextEmbedder(method="sentence_transformers")
# Mock availability to False for the next switch
with patch('semantica.embeddings.text_embedder.FASTEMBED_AVAILABLE', False):
embedder.set_model(method="fastembed", model_name="some-model")
self.assertEqual(embedder.get_method(), "fallback")
self.assertIsNone(embedder.fastembed_model)
self.assertIsNone(embedder.model)
def test_dimension_update_on_switch(self):
"""Test that embedding_dimension is correctly updated when switching models."""
# 1. Setup mocks to return specific dimensions
# FastEmbed mock
mock_fe_instance = MagicMock()
# FastEmbed doesn't have a direct dimension attribute in our implementation,
# it calls _embed_with_fastembed which calls self.fastembed_model.embed
mock_fe_instance.embed.return_value = [np.zeros(384)]
self.mock_fe_class.return_value = mock_fe_instance
# SentenceTransformer mock
mock_st_instance = MagicMock()
mock_st_instance.get_sentence_embedding_dimension.return_value = 768
self.mock_st_class.return_value = mock_st_instance
# 2. Start with FastEmbed (dim 384)
embedder = TextEmbedder(method="fastembed", model_name="fast-model")
self.assertEqual(embedder.get_embedding_dimension(), 384)
# 3. Switch to SentenceTransformer (dim 768)
embedder.set_model(method="sentence_transformers", model_name="st-model-large")
self.assertEqual(embedder.get_embedding_dimension(), 768)
# 4. Switch to another SentenceTransformer with different dimension
mock_st_instance_small = MagicMock()
mock_st_instance_small.get_sentence_embedding_dimension.return_value = 512
self.mock_st_class.return_value = mock_st_instance_small
embedder.set_model(method="sentence_transformers", model_name="st-model-small")
self.assertEqual(embedder.get_embedding_dimension(), 512)
if __name__ == '__main__':
unittest.main()