mirror of
https://github.com/semantica-agi/semantica.git
synced 2026-08-29 04:26:20 +00:00
fix: resolve model switching bug and implement intrinsic dimension detection in TextEmbedder
This commit is contained in:
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user