diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index b17df7340a..c3a7da82e2 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -7,7 +7,9 @@ import pkgutil import re import shutil import sys +import threading import traceback +from contextlib import nullcontext from pathlib import Path from typing import Any, Optional from uuid import uuid4 @@ -66,6 +68,9 @@ if sys.platform == 'darwin' and DEVICE_TYPE == 'cpu': except Exception: pass +# Torch MPS inference is not thread-safe and a concurrent call kills the whole process. +MPS_INFERENCE_LOCK = threading.Lock() if DEVICE_TYPE == 'mps' else nullcontext() + #################################### # LOGGING #################################### diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 294309984e..ec7c73770e 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -32,6 +32,7 @@ from open_webui.env import ( BYPASS_RETRIEVAL_ACCESS_CONTROL, ENABLE_FORWARD_USER_INFO_HEADERS, ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS, + MPS_INFERENCE_LOCK, OFFLINE_MODE, ) from open_webui.models.access_grants import AccessGrants @@ -1116,17 +1117,16 @@ def get_embedding_function( 'SentenceTransformer model name, or configure an external ' 'RAG_EMBEDDING_ENGINE (ollama, openai, azure_openai).' ) - return await asyncio.to_thread( - ( - lambda query, prefix=None: embedding_function.encode( + + def encode(): + with MPS_INFERENCE_LOCK: + return embedding_function.encode( query, batch_size=int(embedding_batch_size), **({'prompt': prefix} if prefix else {}), ).tolist() - ), - query, - prefix, - ) + + return await asyncio.to_thread(encode) return async_embedding_function elif embedding_engine in ['ollama', 'openai', 'azure_openai']: @@ -1250,9 +1250,14 @@ def get_reranking_function(reranking_engine, reranking_model, reranking_function [(query, doc.page_content) for doc in documents], user=user ) else: - return lambda query, documents, user=None: reranking_function.predict( - [(query, doc.page_content) for doc in documents], batch_size=int(reranking_batch_size) - ) + + def predict(query, documents, user=None): + with MPS_INFERENCE_LOCK: + return reranking_function.predict( + [(query, doc.page_content) for doc in documents], batch_size=int(reranking_batch_size) + ) + + return predict # UUIDs, SHA-256 digests, and prefixed variants thereof all fit [A-Za-z0-9_-]. diff --git a/backend/open_webui/routers/evaluations.py b/backend/open_webui/routers/evaluations.py index 8a6bf71f4a..e1c7543eea 100644 --- a/backend/open_webui/routers/evaluations.py +++ b/backend/open_webui/routers/evaluations.py @@ -4,6 +4,7 @@ from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.concurrency import run_in_threadpool from open_webui.constants import ERROR_MESSAGES +from open_webui.env import MPS_INFERENCE_LOCK from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.config import Config @@ -179,8 +180,9 @@ def _compute_similarities(feedbacks: list[LeaderboardFeedbackData], query: str) return {} try: - tag_embeddings = embedding_model.encode(all_tags) - query_embedding = embedding_model.encode([query])[0] + with MPS_INFERENCE_LOCK: + tag_embeddings = embedding_model.encode(all_tags) + query_embedding = embedding_model.encode([query])[0] except Exception as e: log.error(f'Embedding error: {e}') return {}