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
|
|
@ -14,6 +14,11 @@ from utils.rag.config import RAG_EMBED_BATCH_SIZE, RAG_EMBEDDING_MODEL
|
|||
logger = logging.getLogger(__name__)
|
||||
|
||||
_lock = threading.Lock()
|
||||
# Serializes actual compute (encode/tokenize). The ingestion thread and a
|
||||
# concurrent retrieval request share one model; HF's fast tokenizer panics
|
||||
# ("Already borrowed") if encode/tokenize run on it from two threads at once.
|
||||
# GPU work is serial anyway, so this costs nothing in practice.
|
||||
_compute_lock = threading.Lock()
|
||||
_model: Any | None = None
|
||||
_model_name: str | None = None
|
||||
_embedding_dim: int | None = None
|
||||
|
|
@ -66,13 +71,14 @@ def encode(
|
|||
normalize: bool = True,
|
||||
):
|
||||
model = get_embedder(model_name)
|
||||
return model.encode(
|
||||
texts,
|
||||
batch_size = batch_size or RAG_EMBED_BATCH_SIZE,
|
||||
normalize_embeddings = normalize,
|
||||
convert_to_numpy = True,
|
||||
show_progress_bar = False,
|
||||
)
|
||||
with _compute_lock:
|
||||
return model.encode(
|
||||
texts,
|
||||
batch_size = batch_size or RAG_EMBED_BATCH_SIZE,
|
||||
normalize_embeddings = normalize,
|
||||
convert_to_numpy = True,
|
||||
show_progress_bar = False,
|
||||
)
|
||||
|
||||
|
||||
def token_counter(model_name: str | None = None):
|
||||
|
|
@ -81,7 +87,8 @@ def token_counter(model_name: str | None = None):
|
|||
|
||||
def _count(text: str) -> int:
|
||||
try:
|
||||
tokens = model.tokenize([text])
|
||||
with _compute_lock:
|
||||
tokens = model.tokenize([text])
|
||||
ids = tokens.get("input_ids")
|
||||
if ids is None:
|
||||
return max(1, len(text) // 4)
|
||||
|
|
|
|||
|
|
@ -3,20 +3,20 @@
|
|||
|
||||
"""RAG ingestion pipeline.
|
||||
|
||||
Spawn-subprocess per job (parse/chunk/embed); parent persists chunks,
|
||||
vectors, and rebuilds BM25 on completion. Only the parent opens rag.db.
|
||||
In-process daemon thread per job (parse/chunk/embed); the pump thread persists
|
||||
chunks, vectors, and the FTS5 index. The worker shares the warm embedder
|
||||
singleton, so the model is not reloaded per upload.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import queue as queue_module
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Callable
|
||||
from uuid import uuid4
|
||||
|
||||
from loggers import get_logger
|
||||
|
|
@ -33,11 +33,10 @@ from .vector_store import kb_scope, thread_scope
|
|||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_CTX = mp.get_context("spawn")
|
||||
_QUEUE_TIMEOUT_SECONDS = 300
|
||||
|
||||
|
||||
# --- Subprocess worker ---
|
||||
# --- Ingestion worker (runs in a daemon thread) ---
|
||||
|
||||
_MIME_TO_EXT = {
|
||||
"image/png": ".png",
|
||||
|
|
@ -51,30 +50,18 @@ _MIME_TO_EXT = {
|
|||
}
|
||||
|
||||
|
||||
def _subprocess_worker(
|
||||
def _worker(
|
||||
stored_path: str,
|
||||
model_name: str,
|
||||
chunk_size: int,
|
||||
overlap: int,
|
||||
batch_size: int,
|
||||
out_queue: Any,
|
||||
should_cancel: Callable[[], bool],
|
||||
vlm_url: str | None = None,
|
||||
vlm_model: str | None = None,
|
||||
enable_captions: bool = True,
|
||||
) -> None:
|
||||
# Spawned subprocess: structlog setup only ran in the parent's FastAPI
|
||||
# process. Configure it here too so captioner/parser logs render as JSON,
|
||||
# not structlog's default dev ConsoleRenderer.
|
||||
try:
|
||||
import os as _os
|
||||
|
||||
from loggers.config import LogConfig
|
||||
|
||||
LogConfig.setup_logging(
|
||||
env = _os.getenv("ENVIRONMENT_TYPE", "production"),
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
try:
|
||||
from core.rag.captioner import caption_images
|
||||
from core.rag.chunking import chunk_pages
|
||||
|
|
@ -142,11 +129,14 @@ def _subprocess_worker(
|
|||
model = model,
|
||||
chunk_pages = chunk_pages,
|
||||
out_queue = out_queue,
|
||||
should_cancel = should_cancel,
|
||||
send_complete = False,
|
||||
)
|
||||
if should_cancel():
|
||||
return
|
||||
out_queue.put({"type": "complete", "num_chunks": text_count})
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("ingestion subprocess failed")
|
||||
logger.exception("ingestion worker failed")
|
||||
out_queue.put({"type": "error", "error": f"{type(exc).__name__}: {exc}"})
|
||||
|
||||
|
||||
|
|
@ -161,6 +151,7 @@ def _run_standard_chunking(
|
|||
model,
|
||||
chunk_pages,
|
||||
out_queue,
|
||||
should_cancel: Callable[[], bool] = lambda: False,
|
||||
send_complete: bool = True,
|
||||
) -> int:
|
||||
"""Stream text chunks; returns count. send_complete=False when images follow."""
|
||||
|
|
@ -181,14 +172,17 @@ def _run_standard_chunking(
|
|||
|
||||
total = len(chunks)
|
||||
for i in range(0, total, batch_size):
|
||||
if should_cancel():
|
||||
return i
|
||||
batch = chunks[i : i + batch_size]
|
||||
vectors = model.encode(
|
||||
[c.text for c in batch],
|
||||
batch_size = batch_size,
|
||||
normalize_embeddings = True,
|
||||
convert_to_numpy = True,
|
||||
show_progress_bar = False,
|
||||
)
|
||||
with embeddings._compute_lock:
|
||||
vectors = model.encode(
|
||||
[c.text for c in batch],
|
||||
batch_size = batch_size,
|
||||
normalize_embeddings = True,
|
||||
convert_to_numpy = True,
|
||||
show_progress_bar = False,
|
||||
)
|
||||
out_queue.put(
|
||||
{
|
||||
"type": "chunks_batch",
|
||||
|
|
@ -232,7 +226,7 @@ class _JobState:
|
|||
self.progress: float = 0.0
|
||||
self.error: str | None = None
|
||||
self.cancelled = False
|
||||
self.proc: Any = None
|
||||
self.worker: threading.Thread | None = None
|
||||
self.out_queue: Any = None
|
||||
self.subscribers: list[queue_module.Queue[dict]] = []
|
||||
self.lock = threading.Lock()
|
||||
|
|
@ -429,10 +423,10 @@ def _replace_document_pages(document_id: str, pages: list[dict]) -> None:
|
|||
|
||||
def _pump(
|
||||
state: _JobState,
|
||||
proc: Any,
|
||||
worker: threading.Thread,
|
||||
out_queue: Any,
|
||||
) -> None:
|
||||
"""Drain queue until subprocess completes/errors/dies."""
|
||||
"""Drain queue until the worker thread completes/errors/dies."""
|
||||
text_buffer: list[dict] = []
|
||||
embedding_dim: int | None = None
|
||||
final_status = "failed"
|
||||
|
|
@ -452,8 +446,8 @@ def _pump(
|
|||
try:
|
||||
msg = out_queue.get(timeout = _QUEUE_TIMEOUT_SECONDS)
|
||||
except queue_module.Empty:
|
||||
if not proc.is_alive():
|
||||
final_error = "subprocess exited without completion message"
|
||||
if not worker.is_alive():
|
||||
final_error = "worker exited without completion message"
|
||||
break
|
||||
continue
|
||||
mtype = msg.get("type")
|
||||
|
|
@ -514,10 +508,11 @@ def _pump(
|
|||
else:
|
||||
logger.warning("ingestion: unknown message type", mtype = repr(mtype))
|
||||
finally:
|
||||
proc.join(timeout = 30)
|
||||
if proc.is_alive():
|
||||
proc.terminate()
|
||||
proc.join(timeout = 5)
|
||||
# Cooperative: the worker exits on completion/error or when it observes
|
||||
# state.cancelled at a batch boundary. A wedged worker (e.g. a stuck model
|
||||
# download) can't be force-killed in-process; it's a daemon thread, so it
|
||||
# dies with the process. join() bounds how long we wait on the clean path.
|
||||
worker.join(timeout = 30)
|
||||
|
||||
finished_at = int(time.time())
|
||||
if state.cancelled:
|
||||
|
|
@ -613,7 +608,7 @@ def enqueue_ingestion(
|
|||
embedding_model: str | None = None,
|
||||
enable_captions: bool = True,
|
||||
) -> str:
|
||||
"""Create the job row, spawn the subprocess, start the pump; return job_id."""
|
||||
"""Create the job row, start the worker + pump threads; return job_id."""
|
||||
from utils.rag.config import resolve_embedder
|
||||
|
||||
scope = _scope_for(kb_id, thread_id)
|
||||
|
|
@ -655,10 +650,10 @@ def enqueue_ingestion(
|
|||
with _jobs_lock:
|
||||
_jobs[job_id] = state
|
||||
|
||||
out_queue = _CTX.Queue()
|
||||
out_queue: queue_module.Queue = queue_module.Queue()
|
||||
state.out_queue = out_queue
|
||||
proc = _CTX.Process(
|
||||
target = _subprocess_worker,
|
||||
worker = threading.Thread(
|
||||
target = _worker,
|
||||
args = (
|
||||
str(stored_path),
|
||||
model_name,
|
||||
|
|
@ -666,17 +661,19 @@ def enqueue_ingestion(
|
|||
RAG_CHUNK_OVERLAP,
|
||||
RAG_EMBED_BATCH_SIZE,
|
||||
out_queue,
|
||||
lambda: state.cancelled,
|
||||
vlm_url,
|
||||
vlm_model,
|
||||
enable_captions,
|
||||
),
|
||||
name = f"rag-ingest-{job_id[:8]}",
|
||||
daemon = True,
|
||||
)
|
||||
proc.start()
|
||||
state.proc = proc
|
||||
worker.start()
|
||||
state.worker = worker
|
||||
pump_thread = threading.Thread(
|
||||
target = _pump,
|
||||
args = (state, proc, out_queue),
|
||||
args = (state, worker, out_queue),
|
||||
name = f"rag-ingest-pump-{job_id[:8]}",
|
||||
daemon = True,
|
||||
)
|
||||
|
|
@ -685,10 +682,10 @@ def enqueue_ingestion(
|
|||
|
||||
|
||||
def cancel_ingestion(job_id: str) -> bool:
|
||||
"""Stop an in-flight ingestion: wake the pump via a sentinel and kill the
|
||||
worker subprocess so it stops consuming GPU/CPU. Returns False if the job
|
||||
is unknown or already terminal. Artifact/row cleanup is the caller's job
|
||||
(the route deletes the document)."""
|
||||
"""Stop an in-flight ingestion: set the cancel flag (the worker checks it at
|
||||
each batch boundary and stops consuming GPU/CPU) and wake the pump via a
|
||||
sentinel. Returns False if the job is unknown or already terminal.
|
||||
Artifact/row cleanup is the caller's job (the route deletes the document)."""
|
||||
state = get_job_state(job_id)
|
||||
if state is None:
|
||||
return False
|
||||
|
|
@ -700,9 +697,6 @@ def cancel_ingestion(job_id: str) -> bool:
|
|||
state.out_queue.put_nowait({"type": "__cancel__"})
|
||||
except Exception:
|
||||
pass
|
||||
proc = state.proc
|
||||
if proc is not None and proc.is_alive():
|
||||
proc.terminate()
|
||||
return True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import storage.studio_db as studio_db
|
|||
from core.rag.chunking import chunk_pages, chunk_pages_with_spans
|
||||
from core.rag.ingestion import (
|
||||
_JobState,
|
||||
_insert_chunks_and_collect_for_bm25,
|
||||
_insert_chunks,
|
||||
_pump,
|
||||
_replace_document_pages,
|
||||
)
|
||||
|
|
@ -111,7 +111,7 @@ def test_image_chunk_persistence_keeps_page_focus_and_null_text_locators(
|
|||
(doc_id, kb_id, "image.pdf", "application/pdf", "image.pdf", 1_700_000_001),
|
||||
)
|
||||
|
||||
_insert_chunks_and_collect_for_bm25(
|
||||
_insert_chunks(
|
||||
doc_id,
|
||||
"kb_scope",
|
||||
0,
|
||||
|
|
@ -228,16 +228,13 @@ class _OneMessageQueue:
|
|||
return self.message
|
||||
|
||||
|
||||
class _FinishedProcess:
|
||||
class _FinishedWorker:
|
||||
def join(self, timeout: float | None = None) -> None:
|
||||
return None
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
return False
|
||||
|
||||
def terminate(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"pages",
|
||||
|
|
@ -272,7 +269,7 @@ def test_document_pages_missing_document_fails_pump_cleanly(
|
|||
}
|
||||
)
|
||||
|
||||
_pump(state, _FinishedProcess(), queue)
|
||||
_pump(state, _FinishedWorker(), queue)
|
||||
|
||||
assert state.status == "failed"
|
||||
assert state.error is not None
|
||||
|
|
|
|||
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