From 2bd1d06eb2e834f73ff1b6ec46ef088e3228e549 Mon Sep 17 00:00:00 2001 From: KaifAhmad1 Date: Thu, 8 Jan 2026 18:56:42 +0530 Subject: [PATCH] fix: resolve model switching bug and implement intrinsic dimension detection in TextEmbedder --- semantica/embeddings/text_embedder.py | 60 +++++------ tests/embeddings/test_model_switching.py | 124 +++++++++++++++++++++++ 2 files changed, 155 insertions(+), 29 deletions(-) create mode 100644 tests/embeddings/test_model_switching.py diff --git a/semantica/embeddings/text_embedder.py b/semantica/embeddings/text_embedder.py index 446032d3..584b8e50 100644 --- a/semantica/embeddings/text_embedder.py +++ b/semantica/embeddings/text_embedder.py @@ -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( @@ -165,19 +181,12 @@ class TextEmbedder: "Install with: pip install sentence-transformers. " "Using fallback embedding method." ) - - def get_method(self) -> str: - """Get current embedding method.""" - return self.method - - 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 - } + + # If no model loaded, use fallback dimension + if self.embedding_dimension is None: + self.embedding_dimension = self.config.get("dimension", 128) + + 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: """ diff --git a/tests/embeddings/test_model_switching.py b/tests/embeddings/test_model_switching.py new file mode 100644 index 00000000..69b8c2ae --- /dev/null +++ b/tests/embeddings/test_model_switching.py @@ -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()