feat(memory): Complete fastembed migration — fix vector_search.py bug, update all model names, fix test mocks
Co-Authored-By: @memory <memory@aipass>
This commit is contained in:
@@ -251,14 +251,14 @@ def process_file_to_vectors(
|
||||
# Chunk content
|
||||
chunks = chunk_content(content, chunk_size, chunk_overlap)
|
||||
|
||||
# Import ChromaDB and sentence transformers (late import for venv compatibility)
|
||||
# Import ChromaDB and fastembed (late import for venv compatibility)
|
||||
try:
|
||||
import chromadb
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
client = chromadb.PersistentClient(path=str(CHROMA_PATH))
|
||||
collection = client.get_or_create_collection(name=collection_name)
|
||||
model = SentenceTransformer("all-MiniLM-L6-v2")
|
||||
model = TextEmbedding("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
# Generate embeddings and store
|
||||
documents = []
|
||||
@@ -280,7 +280,7 @@ def process_file_to_vectors(
|
||||
)
|
||||
|
||||
# Batch encode
|
||||
embeddings = model.encode(documents).tolist()
|
||||
embeddings = [e.tolist() for e in model.embed(documents)]
|
||||
|
||||
# Upsert (update if exists, insert if not)
|
||||
collection.upsert(documents=documents, embeddings=embeddings, ids=ids, metadatas=metadatas)
|
||||
|
||||
@@ -24,8 +24,7 @@ Design:
|
||||
|
||||
Dependencies (optional):
|
||||
- chromadb
|
||||
- sentence-transformers
|
||||
- torch
|
||||
- fastembed
|
||||
"""
|
||||
|
||||
from typing import List, Dict, Any
|
||||
@@ -56,7 +55,7 @@ class QueryEncoder:
|
||||
Uses all-MiniLM-L6-v2 model with same settings.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
|
||||
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"):
|
||||
"""
|
||||
Initialize query encoder
|
||||
|
||||
@@ -64,28 +63,17 @@ class QueryEncoder:
|
||||
model_name: HuggingFace model identifier (must match embedder.py)
|
||||
|
||||
Raises:
|
||||
ImportError: If sentence-transformers or torch are not installed
|
||||
ImportError: If fastembed is not installed
|
||||
"""
|
||||
# Late imports (heavy optional dependencies)
|
||||
try:
|
||||
import torch
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError as e:
|
||||
logger.info(f"[vector_search] Optional ML dependencies not available: {e}")
|
||||
raise ImportError(
|
||||
f"Search requires sentence-transformers and torch. "
|
||||
f"Install with: pip install sentence-transformers torch. "
|
||||
f"Original error: {e}"
|
||||
)
|
||||
raise ImportError(f"Search requires fastembed. Install with: pip install fastembed. Original error: {e}")
|
||||
|
||||
self.model_name = model_name
|
||||
self.model = SentenceTransformer(model_name)
|
||||
|
||||
# GPU optimization if available
|
||||
self.use_gpu = torch.cuda.is_available()
|
||||
if self.use_gpu:
|
||||
self.model = self.model.to("cuda")
|
||||
|
||||
self.model = TextEmbedding(model_name)
|
||||
self.dimension = 384 # all-MiniLM-L6-v2 output dimension
|
||||
|
||||
def encode(self, query: str) -> List[float]:
|
||||
@@ -98,21 +86,8 @@ class QueryEncoder:
|
||||
Returns:
|
||||
384-dimensional embedding as list
|
||||
"""
|
||||
import torch
|
||||
|
||||
# Encode with same settings as embedder.py
|
||||
embedding = self.model.encode(
|
||||
query,
|
||||
convert_to_tensor=False, # Return numpy
|
||||
normalize_embeddings=True, # Critical for L2 distance
|
||||
show_progress_bar=False,
|
||||
)
|
||||
|
||||
# Cleanup GPU memory if used
|
||||
if self.use_gpu:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return embedding.tolist()
|
||||
embeddings = list(self.model.embed([query]))
|
||||
return embeddings[0].tolist()
|
||||
|
||||
|
||||
# Global encoder instance (singleton pattern)
|
||||
|
||||
@@ -32,7 +32,7 @@ def main():
|
||||
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
model = TextEmbedding("all-MiniLM-L6-v2")
|
||||
model = TextEmbedding("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
sorted_pairs = sorted(enumerate(texts), key=lambda x: len(x[1]))
|
||||
sorted_indices, sorted_texts = zip(*sorted_pairs)
|
||||
|
||||
@@ -30,7 +30,7 @@ logger = get_system_logger()
|
||||
class EmbeddingService:
|
||||
"""Embedding service using fastembed (ONNX runtime, no torch dependency)."""
|
||||
|
||||
def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
|
||||
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"):
|
||||
try:
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError as e:
|
||||
|
||||
@@ -336,14 +336,14 @@ class TestProcessFileToVectors:
|
||||
mock_chromadb = MagicMock()
|
||||
mock_chromadb.PersistentClient.return_value = mock_client
|
||||
|
||||
# Mock sentence_transformers
|
||||
# Mock fastembed
|
||||
mock_model = MagicMock()
|
||||
mock_model.encode.return_value = MagicMock(tolist=MagicMock(return_value=[[0.1, 0.2]]))
|
||||
mock_st = MagicMock()
|
||||
mock_st.SentenceTransformer.return_value = mock_model
|
||||
mock_model.embed.return_value = iter([MagicMock(tolist=MagicMock(return_value=[0.1, 0.2]))])
|
||||
mock_fastembed = MagicMock()
|
||||
mock_fastembed.TextEmbedding.return_value = mock_model
|
||||
|
||||
monkeypatch.setitem(sys.modules, "chromadb", mock_chromadb)
|
||||
monkeypatch.setitem(sys.modules, "sentence_transformers", mock_st)
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
result = mod.process_file_to_vectors(test_file, "test_collection")
|
||||
|
||||
@@ -368,7 +368,7 @@ class TestProcessFileToVectors:
|
||||
|
||||
# Remove chromadb from modules so the import inside the function fails
|
||||
monkeypatch.delitem(sys.modules, "chromadb", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "sentence_transformers", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "fastembed", raising=False)
|
||||
|
||||
# Patch the builtins __import__ to raise for chromadb
|
||||
original_import = __builtins__.__import__ if hasattr(__builtins__, "__import__") else __import__
|
||||
|
||||
@@ -87,26 +87,22 @@ def _prepare_vector_search_mocks(monkeypatch):
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
# Mock sentence_transformers and torch for QueryEncoder
|
||||
# Mock fastembed for QueryEncoder
|
||||
mock_model = MagicMock()
|
||||
mock_model.encode.return_value = MagicMock(tolist=MagicMock(return_value=[0.1] * 384))
|
||||
mock_model.to.return_value = mock_model
|
||||
mock_model.embed.side_effect = lambda texts: iter(
|
||||
[MagicMock(tolist=MagicMock(return_value=[0.1] * 384)) for _ in texts]
|
||||
)
|
||||
|
||||
mock_st_cls = MagicMock(return_value=mock_model)
|
||||
mock_te_cls = MagicMock(return_value=mock_model)
|
||||
|
||||
mock_sentence_transformers = MagicMock()
|
||||
mock_sentence_transformers.SentenceTransformer = mock_st_cls
|
||||
monkeypatch.setitem(sys.modules, "sentence_transformers", mock_sentence_transformers)
|
||||
|
||||
mock_torch = MagicMock()
|
||||
mock_torch.cuda.is_available.return_value = False
|
||||
monkeypatch.setitem(sys.modules, "torch", mock_torch)
|
||||
mock_fastembed = MagicMock()
|
||||
mock_fastembed.TextEmbedding = mock_te_cls
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
return {
|
||||
"client": mock_client,
|
||||
"model": mock_model,
|
||||
"st_cls": mock_st_cls,
|
||||
"torch": mock_torch,
|
||||
"st_cls": mock_te_cls,
|
||||
}
|
||||
|
||||
|
||||
@@ -443,7 +439,7 @@ class TestEncodeQuery:
|
||||
assert result["success"] is True
|
||||
assert len(result["embedding"]) == 384
|
||||
assert result["dimension"] == 384
|
||||
assert result["model"] == "all-MiniLM-L6-v2"
|
||||
assert result["model"] == "sentence-transformers/all-MiniLM-L6-v2"
|
||||
|
||||
def test_encode_rejects_empty_query(self, monkeypatch):
|
||||
"""Empty query string returns error."""
|
||||
@@ -467,7 +463,7 @@ class TestEncodeQuery:
|
||||
"""If model.encode raises, error is caught and returned."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["model"].encode.side_effect = RuntimeError("CUDA out of memory")
|
||||
mocks["model"].embed.side_effect = RuntimeError("CUDA out of memory")
|
||||
|
||||
result = mod.encode_query("test query")
|
||||
|
||||
@@ -481,7 +477,7 @@ class TestEncodeQuery:
|
||||
mod.encode_query("first query")
|
||||
mod.encode_query("second query")
|
||||
|
||||
# SentenceTransformer should only be constructed once
|
||||
# TextEmbedding should only be constructed once
|
||||
mocks["st_cls"].assert_called_once()
|
||||
|
||||
|
||||
|
||||
@@ -236,7 +236,7 @@ class TestPublicGetModelInfo:
|
||||
result = embedder.get_model_info()
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["model_name"] == "all-MiniLM-L6-v2"
|
||||
assert result["model_name"] == "sentence-transformers/all-MiniLM-L6-v2"
|
||||
assert result["dimension"] == 384
|
||||
|
||||
def test_service_init_failure_returns_error(self, monkeypatch):
|
||||
|
||||
Reference in New Issue
Block a user