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:
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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">
|
||||
|
||||
Reference in New Issue
Block a user