Studio RAG: remove cross-encoder reranker (RRF hybrid suffices)
This commit is contained in:
parent
9e75d26293
commit
f955738075
13 changed files with 11 additions and 471 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -37,7 +37,6 @@ export interface PersistedChatSettings {
|
|||
toolCallTimeout?: number;
|
||||
ragSource?: RagSource;
|
||||
ragMode?: RagMode;
|
||||
enableRerank?: boolean;
|
||||
ragTopK?: number;
|
||||
ragMinScore?: number;
|
||||
ragIndexConcurrency?: number;
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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[]> {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue