From 4d2f18981051205016bd24d39521e25a33581225 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Fri, 17 Apr 2026 08:35:45 +0900 Subject: [PATCH] feat: add RAG_RERANKING_BATCH_SIZE configuration option Add configurable reranker batch size (env var RAG_RERANKING_BATCH_SIZE, default 32) following the same pattern as RAG_EMBEDDING_BATCH_SIZE. - config.py: PersistentConfig for RAG_RERANKING_BATCH_SIZE - main.py: import, state init, pass to get_reranking_function - colbert.py: accept batch_size param in predict() (was hardcoded 32) - utils.py: get_reranking_function passes batch_size at call time - retrieval.py: expose in config GET/POST endpoints and ConfigForm - Documents.svelte: add Reranking Batch Size input in admin settings Closes #23730 --- backend/open_webui/config.py | 6 ++++++ backend/open_webui/main.py | 3 +++ backend/open_webui/retrieval/models/colbert.py | 6 +++--- backend/open_webui/retrieval/utils.py | 4 ++-- backend/open_webui/routers/retrieval.py | 9 +++++++++ .../components/admin/Settings/Documents.svelte | 17 +++++++++++++++++ 6 files changed, 40 insertions(+), 5 deletions(-) diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 65f29a4ad..a68720a7c 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -2965,6 +2965,12 @@ RAG_RERANKING_MODEL_TRUST_REMOTE_CODE = ( os.environ.get('RAG_RERANKING_MODEL_TRUST_REMOTE_CODE', 'True').lower() == 'true' ) +RAG_RERANKING_BATCH_SIZE = PersistentConfig( + 'RAG_RERANKING_BATCH_SIZE', + 'rag.reranking_batch_size', + int(os.environ.get('RAG_RERANKING_BATCH_SIZE', '32')), +) + RAG_EXTERNAL_RERANKER_URL = PersistentConfig( 'RAG_EXTERNAL_RERANKER_URL', 'rag.external_reranker_url', diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 63580f990..c183a9323 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -248,6 +248,7 @@ from open_webui.config import ( RAG_EXTERNAL_RERANKER_URL, RAG_EXTERNAL_RERANKER_API_KEY, RAG_EXTERNAL_RERANKER_TIMEOUT, + RAG_RERANKING_BATCH_SIZE, RAG_RERANKING_MODEL_AUTO_UPDATE, RAG_RERANKING_MODEL_TRUST_REMOTE_CODE, RAG_EMBEDDING_ENGINE, @@ -1044,6 +1045,7 @@ app.state.config.RAG_RERANKING_MODEL = RAG_RERANKING_MODEL app.state.config.RAG_EXTERNAL_RERANKER_URL = RAG_EXTERNAL_RERANKER_URL app.state.config.RAG_EXTERNAL_RERANKER_API_KEY = RAG_EXTERNAL_RERANKER_API_KEY app.state.config.RAG_EXTERNAL_RERANKER_TIMEOUT = RAG_EXTERNAL_RERANKER_TIMEOUT +app.state.config.RAG_RERANKING_BATCH_SIZE = RAG_RERANKING_BATCH_SIZE app.state.config.RAG_TEMPLATE = RAG_TEMPLATE @@ -1193,6 +1195,7 @@ app.state.RERANKING_FUNCTION = get_reranking_function( app.state.config.RAG_RERANKING_ENGINE, app.state.config.RAG_RERANKING_MODEL, reranking_function=app.state.rf, + reranking_batch_size=app.state.config.RAG_RERANKING_BATCH_SIZE, ) ######################################## diff --git a/backend/open_webui/retrieval/models/colbert.py b/backend/open_webui/retrieval/models/colbert.py index d122291ec..ceb41824e 100644 --- a/backend/open_webui/retrieval/models/colbert.py +++ b/backend/open_webui/retrieval/models/colbert.py @@ -59,14 +59,14 @@ class ColBERT(BaseReranker): return normalized_scores.detach().cpu().numpy().astype(np.float32) - def predict(self, sentences): + def predict(self, sentences, batch_size=32): query = sentences[0][0] docs = [i[1] for i in sentences] # Embedding the documents - embedded_docs = self.ckpt.docFromText(docs, bsize=32)[0] + embedded_docs = self.ckpt.docFromText(docs, bsize=batch_size)[0] # Embedding the queries - embedded_queries = self.ckpt.queryFromText([query], bsize=32) + embedded_queries = self.ckpt.queryFromText([query], bsize=batch_size) embedded_query = embedded_queries[0] # Calculate retrieval scores for the query against all documents diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 3638409ee..35a72f4df 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -919,7 +919,7 @@ async def generate_embeddings( return embeddings[0] if isinstance(text, str) else embeddings -def get_reranking_function(reranking_engine, reranking_model, reranking_function): +def get_reranking_function(reranking_engine, reranking_model, reranking_function, reranking_batch_size=32): if reranking_function is None: return None if reranking_engine == 'external': @@ -928,7 +928,7 @@ def get_reranking_function(reranking_engine, reranking_model, reranking_function ) else: return lambda query, documents, user=None: reranking_function.predict( - [(query, doc.page_content) for doc in documents] + [(query, doc.page_content) for doc in documents], batch_size=int(reranking_batch_size) ) diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 0c3012202..fd291a280 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -487,6 +487,7 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)): # Reranking settings 'RAG_RERANKING_MODEL': request.app.state.config.RAG_RERANKING_MODEL, 'RAG_RERANKING_ENGINE': request.app.state.config.RAG_RERANKING_ENGINE, + 'RAG_RERANKING_BATCH_SIZE': request.app.state.config.RAG_RERANKING_BATCH_SIZE, 'RAG_EXTERNAL_RERANKER_URL': request.app.state.config.RAG_EXTERNAL_RERANKER_URL, 'RAG_EXTERNAL_RERANKER_API_KEY': request.app.state.config.RAG_EXTERNAL_RERANKER_API_KEY, 'RAG_EXTERNAL_RERANKER_TIMEOUT': request.app.state.config.RAG_EXTERNAL_RERANKER_TIMEOUT, @@ -694,6 +695,7 @@ class ConfigForm(BaseModel): # Reranking settings RAG_RERANKING_MODEL: Optional[str] = None RAG_RERANKING_ENGINE: Optional[str] = None + RAG_RERANKING_BATCH_SIZE: Optional[int] = None RAG_EXTERNAL_RERANKER_URL: Optional[str] = None RAG_EXTERNAL_RERANKER_API_KEY: Optional[str] = None RAG_EXTERNAL_RERANKER_TIMEOUT: Optional[str] = None @@ -940,6 +942,12 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend else request.app.state.config.RAG_EXTERNAL_RERANKER_TIMEOUT ) + request.app.state.config.RAG_RERANKING_BATCH_SIZE = ( + form_data.RAG_RERANKING_BATCH_SIZE + if form_data.RAG_RERANKING_BATCH_SIZE is not None + else request.app.state.config.RAG_RERANKING_BATCH_SIZE + ) + log.info( f'Updating reranking model: {request.app.state.config.RAG_RERANKING_MODEL} to {form_data.RAG_RERANKING_MODEL}' ) @@ -967,6 +975,7 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend request.app.state.config.RAG_RERANKING_ENGINE, request.app.state.config.RAG_RERANKING_MODEL, request.app.state.rf, + reranking_batch_size=request.app.state.config.RAG_RERANKING_BATCH_SIZE, ) except Exception as e: log.error(f'Error loading reranking model: {e}') diff --git a/src/lib/components/admin/Settings/Documents.svelte b/src/lib/components/admin/Settings/Documents.svelte index ece64afd5..eeb6b18b1 100644 --- a/src/lib/components/admin/Settings/Documents.svelte +++ b/src/lib/components/admin/Settings/Documents.svelte @@ -1185,6 +1185,23 @@ {/if} +
+
+ {$i18n.t('Reranking Batch Size')} +
+ +
+ +
+
+
{$i18n.t('Top K')}