Studio RAG: remove cross-encoder reranker (RRF hybrid suffices)

This commit is contained in:
Roland Tannous 2026-06-02 11:17:25 +04:00
commit f955738075
13 changed files with 11 additions and 471 deletions

View file

@ -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,

View file

@ -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

View file

@ -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 "<default>",
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

View file

@ -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."
),
)

View file

@ -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 "<default>",
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:

View file

@ -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"}
)

View file

@ -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,

View file

@ -37,7 +37,6 @@ export interface PersistedChatSettings {
toolCallTimeout?: number;
ragSource?: RagSource;
ragMode?: RagMode;
enableRerank?: boolean;
ragTopK?: number;
ragMinScore?: number;
ragIndexConcurrency?: number;

View file

@ -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.
</p>
</div>
<div className="flex items-center justify-between">
<div className="flex flex-col">
<span className="text-[13px] font-medium">
Use reranker
</span>
<span className="text-[11px] text-muted-foreground">
Slower; uses GPU. Improves quality for fact-heavy
questions.
</span>
</div>
<Switch
checked={enableRerank}
onCheckedChange={(next) => {
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}
/>
</div>
<div className="flex flex-col gap-1.5">
<div className="flex items-center justify-between">
<label className="text-[12px] font-medium text-muted-foreground">

View file

@ -352,7 +352,6 @@ type ChatRuntimeStore = {
activeNativePathToken: string | null;
ragSource: RagSource;
ragMode: RagMode;
enableRerank: boolean;
ragTopK: number;
// Cosine floor; 0 disables. Set > 0 to drop off-topic hits.
ragMinScore: number;
@ -421,7 +420,6 @@ type ChatRuntimeStore = {
setContextUsage: (usage: ChatRuntimeStore["contextUsage"]) => void;
setRagSource: (source: RagSource) => void;
setRagMode: (mode: RagMode) => void;
setEnableRerank: (value: boolean) => void;
setRagTopK: (value: number) => void;
setRagMinScore: (value: number) => void;
setRagIndexConcurrency: (value: number) => void;
@ -447,7 +445,6 @@ type ScalarSettingKey =
| "toolCallTimeout"
| "ragSource"
| "ragMode"
| "enableRerank"
| "ragTopK"
| "ragMinScore"
| "ragIndexConcurrency"
@ -490,7 +487,6 @@ const SCALAR_SETTING_KEYS = [
"toolCallTimeout",
"ragSource",
"ragMode",
"enableRerank",
"ragTopK",
"ragMinScore",
"ragIndexConcurrency",
@ -714,7 +710,6 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
activeNativePathToken: null,
ragSource: { kind: "thread" },
ragMode: "hybrid",
enableRerank: false,
ragTopK: 5,
ragMinScore: 0,
ragIndexConcurrency: 1,
@ -977,15 +972,6 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
setScalarSettingVersion("ragMode", ragMode, state.ragMode);
return { ragMode };
}),
setEnableRerank: (enableRerank) =>
set((state) => {
setScalarSettingVersion(
"enableRerank",
enableRerank,
state.enableRerank,
);
return { enableRerank };
}),
setRagTopK: (ragTopK) =>
set((state) => {
setScalarSettingVersion("ragTopK", ragTopK, state.ragTopK);

View file

@ -127,8 +127,6 @@ export interface SearchRequest {
top_k?: number;
mode?: "bm25" | "dense" | "hybrid";
document_ids?: string[];
enable_rerank?: boolean;
reranker_model?: string;
/** Cosine-similarity floor (0..1). Hits below are dropped server-side. */
min_score?: number;
}
@ -401,22 +399,6 @@ export async function warmupRagEmbedder(): Promise<void> {
await authFetch("/api/rag/warmup", { method: "POST" });
}
/** Download reranker weights into the HF cache. ~1.1 GB on first call,
* no-op when cached. Triggered by the reranker toggle so the download
* happens on an explicit action, not the first chat turn. */
export async function precacheRagReranker(): Promise<{
ok: boolean;
model: string;
error?: string;
}> {
const response = await authFetch("/api/rag/reranker/precache", {
method: "POST",
});
return parseJsonOrThrow<{ ok: boolean; model: string; error?: string }>(
response,
);
}
// --- Search ---
export async function search(req: SearchRequest): Promise<SearchHit[]> {

View file

@ -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

View file

@ -197,8 +197,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,
@ -207,7 +205,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"
@ -219,7 +216,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,
}
@ -230,7 +226,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