mirror of
https://github.com/semantica-agi/semantica.git
synced 2026-08-29 04:26:20 +00:00
125 lines
5.7 KiB
Python
125 lines
5.7 KiB
Python
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()
|