Skip to content

Commit 28235d0

Browse files
committed
feat: implement lazy loading for faiss and sentence_transformers in RAG services
1 parent bd925f6 commit 28235d0

4 files changed

Lines changed: 45 additions & 3 deletions

File tree

api/services/rag/embeddings.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,6 @@
1010
from google.genai import types as genai_types
1111
from openai import AsyncOpenAI
1212
from tenacity import retry, stop_after_attempt, wait_exponential
13-
14-
from sentence_transformers import SentenceTransformer
1513
from api.config import get_settings
1614

1715
logger = structlog.get_logger(__name__)
@@ -48,6 +46,8 @@ def reload_clients(self):
4846
elif self.provider == "local_bge":
4947
# Forcing CPU for stability as requested
5048
device = "cpu"
49+
from sentence_transformers import SentenceTransformer
50+
5151
self.local_model = SentenceTransformer(str(self.model), device=device)
5252
logger.info("local_model_loaded", model=self.model, device=device)
5353
else:

api/services/rag/filesystem_rag_store.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,20 @@
1111
from pathlib import Path
1212
from typing import Any, Dict, List, Optional, Tuple
1313

14-
import faiss
1514
import numpy as np
1615
import structlog
1716
from filelock import FileLock
1817
from rank_bm25 import BM25Okapi
1918

2019
logger = structlog.get_logger(__name__)
2120

21+
22+
def _load_faiss() -> Any:
23+
"""Import faiss only when filesystem RAG operations need it."""
24+
import faiss
25+
26+
return faiss
27+
2228
@dataclass
2329
class RetrievalResult:
2430
chunk_id: str
@@ -86,6 +92,7 @@ async def ingest(
8692

8793
# 3. Update FAISS Index
8894
dim = len(embeddings[0]) if embeddings else 768
95+
faiss = _load_faiss()
8996
if paths["vectors"].exists():
9097
index = faiss.read_index(str(paths["vectors"]))
9198
else:
@@ -141,6 +148,7 @@ def load_namespace(self, namespace: str):
141148
with open(paths["chunks"], "r", encoding="utf-8") as f:
142149
chunks = json.load(f)
143150

151+
faiss = _load_faiss()
144152
index = faiss.read_index(str(paths["vectors"]))
145153

146154
with open(paths["bm25"], "rb") as f:
@@ -189,6 +197,7 @@ def _bootstrap_indices_from_chunks(self, namespace: str) -> None:
189197
return
190198

191199
dim = len(embeddings[0])
200+
faiss = _load_faiss()
192201
index = faiss.IndexHNSWFlat(dim, 32)
193202
index.hnsw.efConstruction = 200
194203
vectors = np.array(embeddings).astype("float32")
@@ -242,6 +251,7 @@ async def retrieve(
242251
bm25 = data["bm25"]
243252

244253
# 1. Vector Search (FAISS)
254+
faiss = _load_faiss()
245255
xq = np.array([query_embedding]).astype("float32")
246256
faiss.normalize_L2(xq) # Assuming cosine similarity if index is inner product,
247257
# but IndexHNSWFlat uses L2 distance by default.

api/tests/test_lazy_imports.py

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,31 @@
1+
import builtins
2+
import importlib
3+
import sys
4+
5+
6+
def test_embeddings_module_import_does_not_require_sentence_transformers(monkeypatch):
7+
original_import = builtins.__import__
8+
9+
def guarded_import(name, globals=None, locals=None, fromlist=(), level=0):
10+
if name == "sentence_transformers":
11+
raise ModuleNotFoundError("blocked for test")
12+
return original_import(name, globals, locals, fromlist, level)
13+
14+
monkeypatch.setattr(builtins, "__import__", guarded_import)
15+
sys.modules.pop("api.services.rag.embeddings", None)
16+
module = importlib.import_module("api.services.rag.embeddings")
17+
assert module is not None
18+
19+
20+
def test_filesystem_rag_store_module_import_does_not_require_faiss(monkeypatch):
21+
original_import = builtins.__import__
22+
23+
def guarded_import(name, globals=None, locals=None, fromlist=(), level=0):
24+
if name == "faiss":
25+
raise ModuleNotFoundError("blocked for test")
26+
return original_import(name, globals, locals, fromlist, level)
27+
28+
monkeypatch.setattr(builtins, "__import__", guarded_import)
29+
sys.modules.pop("api.services.rag.filesystem_rag_store", None)
30+
module = importlib.import_module("api.services.rag.filesystem_rag_store")
31+
assert module is not None

pyproject.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ dependencies = [
2525
"sentry-sdk[fastapi]>=2.20.0",
2626
"slowapi>=0.1.9",
2727
"standardwebhooks>=1.0.0",
28+
"jsonschema>=4.0.0",
2829
"tiktoken>=0.7.0",
2930
"faiss-cpu>=1.8.0",
3031
"rank-bm25>=0.2.2",

0 commit comments

Comments
 (0)