Studio: add BM25/semantic/hybrid search-mode toggle to RAG settings
This commit is contained in:
parent
e48d836ffc
commit
3b477a816c
7 changed files with 74 additions and 14 deletions
|
|
@ -538,7 +538,7 @@ def execute_tool(
|
|||
``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?}`` — consumed by ``search_knowledge_base``.
|
||||
reranker_model?, min_score?, mode?}`` — consumed by ``search_knowledge_base``.
|
||||
"""
|
||||
logger.info(
|
||||
f"execute_tool: name={name}, session_id={session_id}, timeout={timeout}"
|
||||
|
|
@ -562,6 +562,8 @@ def execute_tool(
|
|||
from core.rag.tool import search_knowledge_base
|
||||
|
||||
scope = (tool_context or {}).get("rag_scope") or {}
|
||||
raw_mode = scope.get("mode")
|
||||
mode = raw_mode if raw_mode in ("bm25", "dense", "hybrid") else "hybrid"
|
||||
return search_knowledge_base(
|
||||
query = arguments.get("query", ""),
|
||||
top_k = arguments.get("top_k"),
|
||||
|
|
@ -571,6 +573,7 @@ def execute_tool(
|
|||
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,
|
||||
)
|
||||
return f"Unknown tool: {name}"
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ LLM doesn't need to know about KB UUIDs.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from typing import Any, Literal
|
||||
|
||||
from loggers import get_logger
|
||||
|
||||
|
|
@ -93,6 +93,7 @@ def search_knowledge_base(
|
|||
reranker_model: str | None = None,
|
||||
default_top_k: int = 5,
|
||||
min_score: float = 0.0,
|
||||
mode: Literal["bm25", "dense", "hybrid"] = "hybrid",
|
||||
) -> str:
|
||||
"""Execute the RAG search and return a tool-result string.
|
||||
|
||||
|
|
@ -130,9 +131,10 @@ def search_knowledge_base(
|
|||
scope_embedder = resolve_scope_embedder(scope)
|
||||
|
||||
logger.info(
|
||||
"search_knowledge_base: scope=%s embedder=%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 rerank=%s query=%r",
|
||||
scope,
|
||||
scope_embedder or "<default>",
|
||||
mode,
|
||||
k,
|
||||
min_score,
|
||||
enable_rerank,
|
||||
|
|
@ -140,12 +142,22 @@ def search_knowledge_base(
|
|||
)
|
||||
|
||||
try:
|
||||
hits = retrieval.retrieve_hybrid(
|
||||
scope,
|
||||
query.strip(),
|
||||
k = candidate_k,
|
||||
embedder_model = scope_embedder,
|
||||
)
|
||||
if mode == "bm25":
|
||||
hits = retrieval.retrieve_bm25(scope, query.strip(), candidate_k)
|
||||
elif mode == "dense":
|
||||
hits = retrieval.retrieve_dense(
|
||||
scope,
|
||||
query.strip(),
|
||||
candidate_k,
|
||||
embedder_model = scope_embedder,
|
||||
)
|
||||
else:
|
||||
hits = retrieval.retrieve_hybrid(
|
||||
scope,
|
||||
query.strip(),
|
||||
k = candidate_k,
|
||||
embedder_model = scope_embedder,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("search_knowledge_base retrieval failed")
|
||||
return f"Error: retrieval failed ({type(exc).__name__})."
|
||||
|
|
|
|||
|
|
@ -692,7 +692,8 @@ class ChatCompletionRequest(BaseModel):
|
|||
"[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}. Ignored unless "
|
||||
"default_top_k?: int, reranker_model?: str, min_score?: float, "
|
||||
"mode?: 'bm25'|'dense'|'hybrid'}. Ignored unless "
|
||||
"'search_knowledge_base' is in enabled_tools."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -61,7 +61,7 @@ import {
|
|||
type SearchRequest,
|
||||
search as ragSearch,
|
||||
} from "@/features/rag/api/rag-api";
|
||||
import type { RagSource } from "./chat-settings-api";
|
||||
import type { RagMode, RagSource } from "./chat-settings-api";
|
||||
import {
|
||||
createOpenAIContainer,
|
||||
listOpenAIContainers,
|
||||
|
|
@ -95,11 +95,12 @@ function buildRagRequest(
|
|||
enableRerank: boolean,
|
||||
topK: number,
|
||||
minScore: number,
|
||||
mode: RagMode,
|
||||
): SearchRequest | null {
|
||||
const base: SearchRequest = {
|
||||
query,
|
||||
top_k: topK,
|
||||
mode: "hybrid",
|
||||
mode,
|
||||
enable_rerank: enableRerank,
|
||||
min_score: minScore,
|
||||
};
|
||||
|
|
@ -1023,6 +1024,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
runtime.enableRerank,
|
||||
runtime.ragTopK,
|
||||
runtime.ragMinScore,
|
||||
runtime.ragMode,
|
||||
);
|
||||
if (ragReq) {
|
||||
try {
|
||||
|
|
@ -1664,6 +1666,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
enable_rerank: runtime.enableRerank,
|
||||
default_top_k: runtime.ragTopK,
|
||||
min_score: runtime.ragMinScore,
|
||||
mode: runtime.ragMode,
|
||||
},
|
||||
}
|
||||
: {}),
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ export type RagSource =
|
|||
| { kind: "thread" }
|
||||
| { kind: "kb"; kbId: string };
|
||||
|
||||
export type RagMode = "bm25" | "dense" | "hybrid";
|
||||
|
||||
export interface PersistedChatSettings {
|
||||
inferenceParams?: PersistedInferenceParams;
|
||||
customPresets?: PersistedChatPreset[];
|
||||
|
|
@ -32,6 +34,7 @@ export interface PersistedChatSettings {
|
|||
maxToolCallsPerMessage?: number;
|
||||
toolCallTimeout?: number;
|
||||
ragSource?: RagSource;
|
||||
ragMode?: RagMode;
|
||||
enableRerank?: boolean;
|
||||
ragTopK?: number;
|
||||
ragMinScore?: number;
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ import {
|
|||
} from "./provider-capabilities";
|
||||
import { useChatRuntimeStore } from "./stores/chat-runtime-store";
|
||||
import type { InferenceParams } from "./types/runtime";
|
||||
import type { RagSource } from "./api/chat-settings-api";
|
||||
import type { RagMode, RagSource } from "./api/chat-settings-api";
|
||||
import { DocumentRow } from "@/features/rag/components/document-row";
|
||||
import { KBCreateDialog } from "@/features/rag/components/kb-create-dialog";
|
||||
import { useKnowledgeBases } from "@/features/rag/hooks/use-knowledge-bases";
|
||||
|
|
@ -441,6 +441,8 @@ export function ChatSettingsPanel({
|
|||
const isGguf = useChatRuntimeStore((s) => s.activeGgufVariant) != null;
|
||||
const ragSource = useChatRuntimeStore((s) => s.ragSource);
|
||||
const setRagSource = useChatRuntimeStore((s) => s.setRagSource);
|
||||
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);
|
||||
|
|
@ -1360,6 +1362,32 @@ export function ChatSettingsPanel({
|
|||
source before sending.
|
||||
</p>
|
||||
</div>
|
||||
<div className="flex flex-col gap-1.5">
|
||||
<label className="text-[12px] font-medium text-muted-foreground">
|
||||
Search mode
|
||||
</label>
|
||||
<Select
|
||||
value={ragMode}
|
||||
onValueChange={(v) => setRagMode(v as RagMode)}
|
||||
disabled={!ragEnabled}
|
||||
>
|
||||
<SelectTrigger className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="hybrid">
|
||||
Hybrid (BM25 + semantic)
|
||||
</SelectItem>
|
||||
<SelectItem value="dense">Semantic only</SelectItem>
|
||||
<SelectItem value="bm25">BM25 (lexical) only</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
<p className="text-[11px] text-muted-foreground">
|
||||
Hybrid blends keyword (BM25) with vector similarity — best
|
||||
default. Semantic-only ignores exact terms; BM25-only ignores
|
||||
meaning.
|
||||
</p>
|
||||
</div>
|
||||
<KBCreateDialog
|
||||
open={kbCreateOpen}
|
||||
onOpenChange={setKbCreateOpen}
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ import {
|
|||
loadChatSettingsWithLegacyImport,
|
||||
savePersistedChatSettingsPatch,
|
||||
} from "../utils/chat-settings-storage";
|
||||
import type { RagSource } from "../api/chat-settings-api";
|
||||
import type { RagMode, RagSource } from "../api/chat-settings-api";
|
||||
|
||||
const HF_TOKEN_KEY = "unsloth_hf_token";
|
||||
export const CHAT_REASONING_ENABLED_KEY = "unsloth_chat_reasoning_enabled";
|
||||
|
|
@ -298,6 +298,7 @@ type ChatRuntimeStore = {
|
|||
modelLoading: boolean;
|
||||
activeNativePathToken: string | null;
|
||||
ragSource: RagSource;
|
||||
ragMode: RagMode;
|
||||
enableRerank: boolean;
|
||||
ragTopK: number;
|
||||
// Cosine-similarity floor for RAG hits. Off (0) by default — set
|
||||
|
|
@ -349,6 +350,7 @@ type ChatRuntimeStore = {
|
|||
clearPendingAudio: () => void;
|
||||
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;
|
||||
|
|
@ -370,6 +372,7 @@ type ScalarSettingKey =
|
|||
| "maxToolCallsPerMessage"
|
||||
| "toolCallTimeout"
|
||||
| "ragSource"
|
||||
| "ragMode"
|
||||
| "enableRerank"
|
||||
| "ragTopK"
|
||||
| "ragMinScore";
|
||||
|
|
@ -407,6 +410,7 @@ const SCALAR_SETTING_KEYS = [
|
|||
"maxToolCallsPerMessage",
|
||||
"toolCallTimeout",
|
||||
"ragSource",
|
||||
"ragMode",
|
||||
"enableRerank",
|
||||
"ragTopK",
|
||||
"ragMinScore",
|
||||
|
|
@ -619,6 +623,7 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
|
|||
modelLoading: false,
|
||||
activeNativePathToken: null,
|
||||
ragSource: { kind: "thread" },
|
||||
ragMode: "hybrid",
|
||||
enableRerank: false,
|
||||
ragTopK: 5,
|
||||
ragMinScore: 0,
|
||||
|
|
@ -841,6 +846,11 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
|
|||
setScalarSettingVersion("ragSource", ragSource, state.ragSource);
|
||||
return { ragSource };
|
||||
}),
|
||||
setRagMode: (ragMode) =>
|
||||
set((state) => {
|
||||
setScalarSettingVersion("ragMode", ragMode, state.ragMode);
|
||||
return { ragMode };
|
||||
}),
|
||||
setEnableRerank: (enableRerank) =>
|
||||
set((state) => {
|
||||
setScalarSettingVersion(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue