RAG: run ingestion in-process thread sharing warm embedder + compute lock
This commit is contained in:
parent
f5673e9bb0
commit
f1d84c09f8
4 changed files with 119 additions and 66 deletions
55
tests/python/test_rag_embeddings.py
Normal file
55
tests/python/test_rag_embeddings.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
"""RAG embedder compute-lock wiring (no GPU / real model needed)."""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
STUDIO_BACKEND = REPO_ROOT / "studio" / "backend"
|
||||
if str(STUDIO_BACKEND) not in sys.path:
|
||||
sys.path.insert(0, str(STUDIO_BACKEND))
|
||||
|
||||
|
||||
class _FakeModel:
|
||||
"""Records whether the shared compute lock was held during each call."""
|
||||
|
||||
def __init__(self, lock) -> None:
|
||||
self._lock = lock
|
||||
self.encode_held: bool | None = None
|
||||
self.tokenize_held: bool | None = None
|
||||
|
||||
def encode(self, texts, **_kw):
|
||||
self.encode_held = self._lock.locked()
|
||||
return np.zeros((len(texts), 3), dtype = "float32")
|
||||
|
||||
def tokenize(self, _texts):
|
||||
self.tokenize_held = self._lock.locked()
|
||||
return {"input_ids": np.zeros((1, 5), dtype = "int64")}
|
||||
|
||||
|
||||
def test_encode_holds_compute_lock_then_releases(monkeypatch):
|
||||
from core.rag import embeddings
|
||||
|
||||
fake = _FakeModel(embeddings._compute_lock)
|
||||
monkeypatch.setattr(embeddings, "get_embedder", lambda *a, **k: fake)
|
||||
|
||||
out = embeddings.encode(["hello"])
|
||||
|
||||
assert fake.encode_held is True
|
||||
assert embeddings._compute_lock.locked() is False
|
||||
assert out.shape == (1, 3)
|
||||
|
||||
|
||||
def test_token_counter_holds_compute_lock_then_releases(monkeypatch):
|
||||
from core.rag import embeddings
|
||||
|
||||
fake = _FakeModel(embeddings._compute_lock)
|
||||
monkeypatch.setattr(embeddings, "get_embedder", lambda *a, **k: fake)
|
||||
|
||||
counter = embeddings.token_counter()
|
||||
n = counter("hello world")
|
||||
|
||||
assert fake.tokenize_held is True
|
||||
assert n == 5
|
||||
assert embeddings._compute_lock.locked() is False
|
||||
Loading…
Add table
Add a link
Reference in a new issue