From f955738075cac6cf74e40e8c0fda7d1fc4ff2944 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 2 Jun 2026 11:17:25 +0400 Subject: [PATCH] Studio RAG: remove cross-encoder reranker (RRF hybrid suffices) --- studio/backend/core/inference/tools.py | 6 +- studio/backend/core/rag/reranker.py | 218 ------------------ studio/backend/core/rag/tool.py | 34 +-- studio/backend/models/inference.py | 5 +- studio/backend/routes/rag.py | 61 +---- studio/backend/utils/rag/config.py | 6 - .../src/features/chat/api/chat-adapter.ts | 1 - .../features/chat/api/chat-settings-api.ts | 1 - .../src/features/chat/chat-settings-sheet.tsx | 49 ---- .../chat/stores/chat-runtime-store.ts | 14 -- .../frontend/src/features/rag/api/rag-api.ts | 18 -- tests/python/test_rag_reranker.py | 64 ----- tests/python/test_rag_tool_handler.py | 5 - 13 files changed, 11 insertions(+), 471 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 b1024c7e86..6ca5d1c17c 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -724,8 +724,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}" @@ -779,8 +779,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..3a3ebb3afa 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", @@ -163,26 +161,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 +233,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/models/inference.py b/studio/backend/models/inference.py index d92756feeb..6ec5cb9a9f 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -744,9 +744,8 @@ class ChatCompletionRequest(BaseModel): description = ( "[x-unsloth] Per-request context the `search_knowledge_base` tool " "consumes when the LLM invokes it. Shape: " - "{kb_id?: str, thread_id?: str, enable_rerank?: bool, " - "default_top_k?: int, reranker_model?: str, min_score?: float, " - "mode?: 'bm25'|'dense'|'hybrid'}. Ignored unless " + "{kb_id?: str, thread_id?: str, default_top_k?: int, " + "min_score?: float, mode?: 'bm25'|'dense'|'hybrid'}. Ignored unless " "'search_knowledge_base' is in enabled_tools." ), ) 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/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 8353fb1f4b..8fb2088214 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -2376,7 +2376,6 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { ragSource.kind === "thread" ? (resolvedThreadId ?? null) : null, - enable_rerank: runtime.enableRerank, default_top_k: runtime.ragTopK, min_score: runtime.ragMinScore, mode: runtime.ragMode, diff --git a/studio/frontend/src/features/chat/api/chat-settings-api.ts b/studio/frontend/src/features/chat/api/chat-settings-api.ts index dc06ad0321..f0c20f30d5 100644 --- a/studio/frontend/src/features/chat/api/chat-settings-api.ts +++ b/studio/frontend/src/features/chat/api/chat-settings-api.ts @@ -37,7 +37,6 @@ export interface PersistedChatSettings { toolCallTimeout?: number; ragSource?: RagSource; ragMode?: RagMode; - enableRerank?: boolean; ragTopK?: number; ragMinScore?: number; ragIndexConcurrency?: number; diff --git a/studio/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx index c75dd25db6..8f28119150 100644 --- a/studio/frontend/src/features/chat/chat-settings-sheet.tsx +++ b/studio/frontend/src/features/chat/chat-settings-sheet.tsx @@ -48,7 +48,6 @@ import { import { type KBMode, type ChunkingStrategy as RagChunkingStrategy, - precacheRagReranker, } from "@/features/rag/api/rag-api"; import { DocumentRow } from "@/features/rag/components/document-row"; import { KBCreateDialog } from "@/features/rag/components/kb-create-dialog"; @@ -523,8 +522,6 @@ export function ChatSettingsPanel({ const ragMode = useChatRuntimeStore((s) => s.ragMode); const setRagMode = useChatRuntimeStore((s) => s.setRagMode); const ragToolEnabled = useChatRuntimeStore((s) => s.ragToolEnabled); - const enableRerank = useChatRuntimeStore((s) => s.enableRerank); - const setEnableRerank = useChatRuntimeStore((s) => s.setEnableRerank); const ragTopK = useChatRuntimeStore((s) => s.ragTopK); const setRagTopK = useChatRuntimeStore((s) => s.setRagTopK); const ragIndexConcurrency = useChatRuntimeStore( @@ -1716,52 +1713,6 @@ export function ChatSettingsPanel({ grounding, more tokens.

-
-
- - Use reranker - - - Slower; uses GPU. Improves quality for fact-heavy - questions. - -
- { - setEnableRerank(next); - if (!next) return; - // First flip-on may download ~1.1 GB; the toast - // covers the latency so the next query doesn't look - // hung waiting on the reranker. - const toastId = toast.loading( - "Preparing reranker (one-time download)…", - ); - void precacheRagReranker() - .then((res) => { - if (res.ok) { - toast.success("Reranker ready", { id: toastId }); - } else { - toast.error( - `Reranker download failed: ${res.error ?? "unknown"}`, - { id: toastId }, - ); - setEnableRerank(false); - } - }) - .catch((err: unknown) => { - toast.error( - `Reranker download failed: ${ - err instanceof Error ? err.message : String(err) - }`, - { id: toastId }, - ); - setEnableRerank(false); - }); - }} - disabled={!ragEnabled} - /> -