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}
- />
-