From b085b37b69cbdd029ca919b4ac930c41a53dd369 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 31 May 2026 14:14:52 +0000 Subject: [PATCH] Studio RAG trim: remove cross-encoder reranker The cross-encoder reranker is off by default (enable_rerank=False everywhere) and adds a second model download plus a candidate-widening pass on every search. Removing it keeps the core retrieval (parse, chunk, embed, BM25 + dense, RRF, search_knowledge_base tool) intact while dropping ~375 lines. - delete core/rag/reranker.py and its test - drop enable_rerank / reranker_model from the search tool, tools dispatch, and the /rag/search route (candidate_k is now just top_k) - remove the /rag/reranker/precache endpoint and reranker config knobs - update tool-handler test to the trimmed rag_scope shape 42 RAG tests pass. --- studio/backend/core/inference/tools.py | 6 +- studio/backend/core/rag/reranker.py | 218 ------------------------- studio/backend/core/rag/tool.py | 33 +--- studio/backend/routes/rag.py | 61 +------ studio/backend/utils/rag/config.py | 6 - tests/python/test_rag_reranker.py | 64 -------- tests/python/test_rag_tool_handler.py | 5 - 7 files changed, 9 insertions(+), 384 deletions(-) delete mode 100644 studio/backend/core/rag/reranker.py delete mode 100644 tests/python/test_rag_reranker.py diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py index 48d70aa67f..00aa3550b1 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -626,8 +626,8 @@ def execute_tool( ``session_id``: optional thread/session ID for per-conversation sandbox isolation. ``tool_context``: optional per-request extras the LLM does not see (RAG scope, future per-tool overrides). Keys consumed: - - ``rag_scope``: ``{kb_id?, thread_id?, enable_rerank?, default_top_k?, - reranker_model?, min_score?, mode?}`` — consumed by ``search_knowledge_base``. + - ``rag_scope``: ``{kb_id?, thread_id?, default_top_k?, min_score?, mode?}`` + — consumed by ``search_knowledge_base``. """ logger.info( f"execute_tool: name={name}, session_id={session_id}, timeout={timeout}" @@ -677,8 +677,6 @@ def execute_tool( top_k = arguments.get("top_k"), scope_kb_id = scope.get("kb_id"), scope_thread_id = scope.get("thread_id"), - enable_rerank = bool(scope.get("enable_rerank")), - reranker_model = scope.get("reranker_model"), default_top_k = int(scope.get("default_top_k") or 5), min_score = float(scope.get("min_score") or 0.0), mode = mode, diff --git a/studio/backend/core/rag/reranker.py b/studio/backend/core/rag/reranker.py deleted file mode 100644 index 2c5bc3ea30..0000000000 --- a/studio/backend/core/rag/reranker.py +++ /dev/null @@ -1,218 +0,0 @@ -# SPDX-License-Identifier: AGPL-3.0-only -# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 - -"""Opt-in CrossEncoder reranker (off by default; shares GPU with chat model).""" - -from __future__ import annotations - -import gc -import sys -import threading -import time -from typing import Any - -from loggers import get_logger -from utils.rag.config import RAG_RERANK_BATCH_SIZE, RAG_RERANKER_MODEL - -from .retrieval import Hit - -logger = get_logger(__name__) - -# Reentrant: get_reranker() holds the lock while calling unload(), which also -# enters `with _lock`. A plain Lock would self-deadlock; RLock allows re-entry. -_lock = threading.RLock() -_model: Any | None = None -_model_name: str | None = None - - -def _resolve_device() -> str: - """Prefer CUDA when available; otherwise CPU. Explicit so we don't rely - on sentence-transformers' auto-detect (which historically picks CPU when - CUDA_VISIBLE_DEVICES is set funny).""" - try: - import torch - - if torch.cuda.is_available(): - return "cuda" - except Exception: # noqa: BLE001 - pass - return "cpu" - - -def _load(model_name: str) -> Any: - # Unconditional stderr print so this shows even when structlog routing - # misbehaves — diagnostics for a previously invisible hang. - print( - f"[rag.reranker] _load entered: model={model_name}", - file = sys.stderr, - flush = True, - ) - from sentence_transformers import CrossEncoder - - device = _resolve_device() - print( - f"[rag.reranker] device resolved: {device}", - file = sys.stderr, - flush = True, - ) - logger.info( - "Loading RAG reranker", - model = model_name, - device = device, - ) - started = time.perf_counter() - print( - f"[rag.reranker] calling CrossEncoder(...) on {device}", - file = sys.stderr, - flush = True, - ) - model = CrossEncoder(model_name, device = device) - elapsed = round(time.perf_counter() - started, 2) - print( - f"[rag.reranker] CrossEncoder returned in {elapsed}s", - file = sys.stderr, - flush = True, - ) - logger.info( - "RAG reranker loaded", - model = model_name, - device = device, - elapsed_seconds = elapsed, - ) - return model - - -def precache_reranker(model_name: str | None = None) -> None: - """Download reranker weights into the HF cache (no instantiation). - - Mirrors ``precache_helper_gguf``: runs in a background thread on - FastAPI startup so the first user-facing rerank doesn't pay the - ~1.1 GB download. Safe to call when the model is already cached - (huggingface_hub no-ops on existing files). - """ - target = model_name or RAG_RERANKER_MODEL - try: - from huggingface_hub import snapshot_download - from huggingface_hub.utils import disable_progress_bars - - disable_progress_bars() - logger.info("Pre-caching RAG reranker", model = target) - started = time.perf_counter() - snapshot_download(repo_id = target, repo_type = "model") - logger.info( - "RAG reranker cached", - model = target, - elapsed_seconds = round(time.perf_counter() - started, 2), - ) - except Exception as exc: # noqa: BLE001 - # Non-critical: the lazy loader retries the download on first use; log it. - logger.warning( - "RAG reranker precache failed; will download lazily", - model = target, - error = str(exc), - ) - - -def get_reranker(model_name: str | None = None) -> Any: - global _model, _model_name - target = model_name or RAG_RERANKER_MODEL - with _lock: - if _model is None or _model_name != target: - unload() - _model = _load(target) - _model_name = target - return _model - - -def unload() -> None: - """Drop the reranker; next call lazy-loads again.""" - global _model, _model_name - with _lock: - if _model is not None: - _model = None - _model_name = None - gc.collect() - try: - import torch - - if torch.cuda.is_available(): - torch.cuda.empty_cache() - except ImportError: - pass - - -def rerank( - query: str, - pairs: list[tuple[Hit, str]], - *, - model_name: str | None = None, - top_k: int | None = None, -) -> list[Hit]: - """Re-order (Hit, text) pairs by CrossEncoder score; image hits are appended last.""" - if not pairs: - return [] - text_pairs = [(h, t) for h, t in pairs if h.kind != "image"] - image_hits = [h for h, _t in pairs if h.kind == "image"] - - print( - f"[rag.reranker] rerank entered: n_pairs={len(text_pairs)}", - file = sys.stderr, - flush = True, - ) - model = get_reranker(model_name) - print( - "[rag.reranker] reranker model in hand", - file = sys.stderr, - flush = True, - ) - if text_pairs: - inputs = [(query, text) for _, text in text_pairs] - print( - f"[rag.reranker] predict starting: n_inputs={len(inputs)} " - f"batch_size={RAG_RERANK_BATCH_SIZE}", - file = sys.stderr, - flush = True, - ) - logger.info( - "RAG reranker predict starting", - n_inputs = len(inputs), - batch_size = RAG_RERANK_BATCH_SIZE, - ) - started = time.perf_counter() - scores = model.predict( - inputs, - batch_size = RAG_RERANK_BATCH_SIZE, - show_progress_bar = False, - ) - elapsed = round(time.perf_counter() - started, 2) - print( - f"[rag.reranker] predict done in {elapsed}s", - file = sys.stderr, - flush = True, - ) - logger.info( - "RAG reranker predict done", - n_inputs = len(inputs), - elapsed_seconds = elapsed, - ) - ranked = sorted( - zip(text_pairs, scores), - key = lambda item: float(item[1]), - reverse = True, - ) - reranked_text = [ - Hit( - chunk_id = h.chunk_id, - score = float(s), - document_id = h.document_id, - chunk_index = h.chunk_index, - kind = h.kind, - ) - for (h, _t), s in ranked - ] - else: - reranked_text = [] - out: list[Hit] = reranked_text + image_hits - if top_k is not None: - out = out[:top_k] - return out diff --git a/studio/backend/core/rag/tool.py b/studio/backend/core/rag/tool.py index c2363b8156..aba448471a 100644 --- a/studio/backend/core/rag/tool.py +++ b/studio/backend/core/rag/tool.py @@ -140,8 +140,6 @@ def search_knowledge_base( top_k: int | None = None, scope_kb_id: str | None = None, scope_thread_id: str | None = None, - enable_rerank: bool = False, - reranker_model: str | None = None, default_top_k: int = 5, min_score: float = 0.0, mode: Literal["bm25", "dense", "hybrid"] = "hybrid", @@ -164,25 +162,19 @@ def search_knowledge_base( scope = kb_scope(scope_kb_id) if scope_kb_id else thread_scope(scope_thread_id) k = top_k if top_k is not None else default_top_k - if enable_rerank: - from utils.rag.config import RAG_RERANK_CANDIDATE_K - - candidate_k = max(k, RAG_RERANK_CANDIDATE_K) - else: - candidate_k = k + candidate_k = k from core.rag.scope import resolve_scope_embedder scope_embedder = resolve_scope_embedder(scope) logger.info( - "search_knowledge_base: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r", + "search_knowledge_base: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f query=%r", scope, scope_embedder or "", mode, k, min_score, - enable_rerank, query[:120], ) @@ -242,26 +234,7 @@ def search_knowledge_base( for row in rows: lookup[row["chunk_id"]] = dict(row) - if enable_rerank and hits: - from core.rag import reranker - - pairs = [ - (hit, lookup[hit.chunk_id]["text"]) - for hit in hits - if hit.chunk_id in lookup - ] - try: - hits = reranker.rerank( - query.strip(), - pairs, - model_name = reranker_model, - top_k = k, - ) - except Exception as exc: # noqa: BLE001 - logger.warning("rerank failed in search_knowledge_base: %s", exc) - hits = hits[:k] - else: - hits = hits[:k] + hits = hits[:k] # Merge Hit metadata (score, dense_score, chunk_index) into the sqlite row so # the formatter sees one flat dict per chunk. Image-kind hits flow through so diff --git a/studio/backend/routes/rag.py b/studio/backend/routes/rag.py index e6d8eaed99..37c2d232a9 100644 --- a/studio/backend/routes/rag.py +++ b/studio/backend/routes/rag.py @@ -41,7 +41,7 @@ async def _sse_auth( return await get_current_subject_sse(token, authorization) -from core.rag import embeddings, ingestion, reranker, retrieval, vector_store +from core.rag import embeddings, ingestion, retrieval, vector_store from core.rag.authorization import document_for_subject_or_404 from core.rag.locators import backfill_document_locators from core.rag.vector_store import kb_scope, thread_scope @@ -54,7 +54,6 @@ from storage.studio_db import ( from utils.paths.storage_roots import ensure_dir, rag_uploads_root, resolve_under_root from utils.rag.config import ( RAG_MAX_UPLOAD_MB, - RAG_RERANK_CANDIDATE_K, RAG_UPLOAD_EXTS, ) @@ -133,8 +132,6 @@ class SearchRequest(BaseModel): top_k: int = Field(default = 10, ge = 1, le = 100) mode: Literal["bm25", "dense", "hybrid"] = "hybrid" document_ids: list[str] | None = None - enable_rerank: bool = False - reranker_model: str | None = None min_score: float = Field(default = 0.0, ge = 0.0, le = 1.0) @@ -508,37 +505,6 @@ def warmup_rag_embedder( return {"ok": True, "model": model_name} -@router.post("/reranker/precache") -def precache_rag_reranker( - current_subject: str = Depends(get_current_subject), -) -> dict: - """Download the reranker weights (~1.1 GB) into the HF cache. - - Called from the frontend the moment the user flips the "Use - reranker" switch ON so the cost lands on the explicit toggle - instead of the first chat turn — where a multi-minute download - looks like a hung tool call. - """ - from core.rag.reranker import precache_reranker - from utils.rag.config import RAG_RERANKER_MODEL - - try: - precache_reranker() - except Exception as exc: # noqa: BLE001 - # Log details server-side; client gets a generic message so paths/stack stay hidden. - logger.warning( - "RAG reranker precache failed", - model = RAG_RERANKER_MODEL, - error = str(exc), - ) - return { - "ok": False, - "model": RAG_RERANKER_MODEL, - "error": "Failed to download reranker", - } - return {"ok": True, "model": RAG_RERANKER_MODEL} - - @router.put("/defaults", response_model = RagDefaults) def set_rag_defaults( payload: UpdateRagDefaultsRequest, @@ -1634,22 +1600,16 @@ def search( # Query must use the same embedder as the scope (dim must match). scope_embedder = _resolve_scope_embedder(scope) logger.info( - "RAG search: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r", + "RAG search: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f query=%r", scope, scope_embedder or "", payload.mode, payload.top_k, payload.min_score, - payload.enable_rerank, payload.query[:120], ) - # Reranker needs a wider candidate pool than top_k. - candidate_k = ( - max(payload.top_k, RAG_RERANK_CANDIDATE_K) - if payload.enable_rerank - else payload.top_k - ) + candidate_k = payload.top_k if payload.mode == "bm25": hits = retrieval.retrieve_bm25(scope, payload.query, candidate_k) @@ -1703,20 +1663,7 @@ def search( for r in rows: chunk_lookup[r["chunk_id"]] = dict(r) - if payload.enable_rerank: - pairs = [ - (hit, chunk_lookup[hit.chunk_id]["text"]) - for hit in hits - if hit.chunk_id in chunk_lookup - ] - hits = reranker.rerank( - payload.query, - pairs, - model_name = payload.reranker_model, - top_k = payload.top_k, - ) - else: - hits = hits[: payload.top_k] + hits = hits[: payload.top_k] out: list[SearchHit] = [] for hit in hits: diff --git a/studio/backend/utils/rag/config.py b/studio/backend/utils/rag/config.py index 9d97039e9c..278a866613 100644 --- a/studio/backend/utils/rag/config.py +++ b/studio/backend/utils/rag/config.py @@ -71,12 +71,6 @@ RAG_MAX_UPLOAD_MB: int = _env_int("UNSLOTH_RAG_MAX_UPLOAD_MB", 50) RAG_EMBED_BATCH_SIZE: int = _env_int("UNSLOTH_RAG_EMBED_BATCH_SIZE", 32) -RAG_RERANKER_MODEL: str = ( - os.environ.get("UNSLOTH_RAG_RERANKER_MODEL", "").strip() or "BAAI/bge-reranker-base" -) -RAG_RERANK_CANDIDATE_K: int = _env_int("UNSLOTH_RAG_RERANK_CANDIDATE_K", 50) -RAG_RERANK_BATCH_SIZE: int = _env_int("UNSLOTH_RAG_RERANK_BATCH_SIZE", 16) - RAG_UPLOAD_EXTS: frozenset[str] = frozenset( {".pdf", ".txt", ".md", ".markdown", ".docx", ".html", ".htm"} ) diff --git a/tests/python/test_rag_reranker.py b/tests/python/test_rag_reranker.py deleted file mode 100644 index c19f22f1c4..0000000000 --- a/tests/python/test_rag_reranker.py +++ /dev/null @@ -1,64 +0,0 @@ -"""Reranker tests — skipped if sentence_transformers is unavailable. - -These tests load a real CrossEncoder, so they're slow and gated under -the ``server`` marker so a default ``pytest`` run skips them. Force -with ``pytest -m server``. -""" - -import sys -from pathlib import Path - -import pytest - -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)) - -pytest.importorskip("sentence_transformers") - - -def test_rerank_empty_returns_empty(): - from core.rag.reranker import rerank - - assert rerank("anything", []) == [] - - -@pytest.mark.server -def test_rerank_reorders_by_relevance(monkeypatch): - """Hide the relevant chunk at the back of the input and check it bubbles up.""" - monkeypatch.setenv( - "UNSLOTH_RAG_RERANKER_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2" - ) - from core.rag.reranker import rerank, unload - from core.rag.retrieval import Hit - - pairs = [ - (Hit("noise1", 0.0), "Cats are small carnivorous mammals."), - (Hit("noise2", 0.0), "The Eiffel Tower is in Paris, France."), - (Hit("noise3", 0.0), "Python is a programming language."), - ( - Hit("answer", 0.0), - "The speed of light in vacuum is approximately 299792458 meters per second.", - ), - ] - try: - ranked = rerank("How fast does light travel?", pairs, top_k = 2) - assert ranked - assert ranked[0].chunk_id == "answer" - finally: - unload() - - -@pytest.mark.server -def test_unload_clears_singleton(monkeypatch): - monkeypatch.setenv( - "UNSLOTH_RAG_RERANKER_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2" - ) - from core.rag import reranker - from core.rag.retrieval import Hit - - reranker.rerank("q", [(Hit("a", 0.0), "some text")]) - assert reranker._model is not None - reranker.unload() - assert reranker._model is None diff --git a/tests/python/test_rag_tool_handler.py b/tests/python/test_rag_tool_handler.py index c3944b5792..c5d4d2dcd5 100644 --- a/tests/python/test_rag_tool_handler.py +++ b/tests/python/test_rag_tool_handler.py @@ -196,8 +196,6 @@ def test_execute_tool_dispatches_to_search_knowledge_base(): top_k = None, scope_kb_id = None, scope_thread_id = None, - enable_rerank = False, - reranker_model = None, default_top_k = 5, min_score = 0.0, **kwargs, @@ -206,7 +204,6 @@ def test_execute_tool_dispatches_to_search_knowledge_base(): called["top_k"] = top_k called["scope_kb_id"] = scope_kb_id called["scope_thread_id"] = scope_thread_id - called["enable_rerank"] = enable_rerank called["default_top_k"] = default_top_k called["min_score"] = min_score return "stub-result" @@ -218,7 +215,6 @@ def test_execute_tool_dispatches_to_search_knowledge_base(): tool_context = { "rag_scope": { "kb_id": "kb-1", - "enable_rerank": True, "default_top_k": 3, "min_score": 0.35, } @@ -229,7 +225,6 @@ def test_execute_tool_dispatches_to_search_knowledge_base(): assert called["top_k"] == 7 assert called["scope_kb_id"] == "kb-1" assert called["scope_thread_id"] is None - assert called["enable_rerank"] is True assert called["default_top_k"] == 3 assert called["min_score"] == 0.35