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
This commit is contained in:
Timothy Jaeryang Baek
2026-04-17 08:35:45 +09:00
parent 70a6a24f14
commit 4d2f189810
6 changed files with 40 additions and 5 deletions
+6
View File
@@ -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',
+3
View File
@@ -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,
)
########################################
@@ -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
+2 -2
View File
@@ -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)
)
+9
View File
@@ -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}')
@@ -1185,6 +1185,23 @@
</div>
{/if}
<div class=" mb-2.5 flex w-full justify-between">
<div class=" self-center text-xs font-medium">
{$i18n.t('Reranking Batch Size')}
</div>
<div class="">
<input
bind:value={RAGConfig.RAG_RERANKING_BATCH_SIZE}
type="number"
class=" bg-transparent text-center w-14 outline-none"
min="1"
max="16000"
step="1"
/>
</div>
</div>
<div class=" mb-2.5 flex w-full justify-between">
<div class=" self-center text-xs font-medium">{$i18n.t('Top K')}</div>
<div class="flex items-center relative">