Commit 436a6311 authored by Kantz's avatar Kantz
Browse files

optimierung der Embedding zeit

parent 90792cb9
from __future__ import annotations
import time
from threading import Lock
from typing import Any, Dict, Tuple
import app.config as config
from app.deterministic_services.embeddings import BaseEmbeddings, EmbeddingFactory
_EMBEDDER_LOCK = Lock()
_CACHED_EMBEDDER: BaseEmbeddings | None = None
_CACHED_KEY: Tuple[Any, ...] | None = None
def _settings_cache_key(settings: config.EmbeddingSettings) -> Tuple[Any, ...]:
return (
settings.embedding_type,
settings.base_url,
settings.model,
settings.target_dim,
)
def get_embedder() -> tuple[BaseEmbeddings, bool]:
settings = config.get_embedding_settings()
key = _settings_cache_key(settings)
global _CACHED_EMBEDDER
global _CACHED_KEY
with _EMBEDDER_LOCK:
if _CACHED_EMBEDDER is not None and _CACHED_KEY == key:
return _CACHED_EMBEDDER, True
_CACHED_EMBEDDER = EmbeddingFactory.create(settings)
_CACHED_KEY = key
return _CACHED_EMBEDDER, False
def warmup_embedder() -> Dict[str, Any]:
started = time.perf_counter()
embedder, cache_hit = get_embedder()
init_ms = round((time.perf_counter() - started) * 1000, 2)
warm_started = time.perf_counter()
embedder.embed_query("warmup")
warmup_ms = round((time.perf_counter() - warm_started) * 1000, 2)
total_ms = round((time.perf_counter() - started) * 1000, 2)
return {
"cache_hit": cache_hit,
"embedder_init_ms": init_ms,
"embed_query_warmup_ms": warmup_ms,
"total_warmup_ms": total_ms,
}
...@@ -9,11 +9,11 @@ import app.config as config ...@@ -9,11 +9,11 @@ import app.config as config
from app.deterministic_services import ( from app.deterministic_services import (
Source, Source,
context_store, context_store,
embedding_provider,
referenz_decoder, referenz_decoder,
retrieval_store, retrieval_store,
tool_logging, tool_logging,
) )
from app.deterministic_services.embeddings import EmbeddingFactory
T = TypeVar("T") T = TypeVar("T")
...@@ -113,23 +113,40 @@ def extract_user_messages(messages: list[dict]) -> list[str]: ...@@ -113,23 +113,40 @@ def extract_user_messages(messages: list[dict]) -> list[str]:
return user_contents return user_contents
def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]: def retrieve_context(
embedder = EmbeddingFactory.create(config.get_embedding_settings()) query_text: str, pg_url: str | None = None
) -> tuple[List[Source], dict]:
embedder_started = time.perf_counter()
embedder, cache_hit = embedding_provider.get_embedder()
embedder_ms = round((time.perf_counter() - embedder_started) * 1000, 2)
url = pg_url or config.get_postgres_url() url = pg_url or config.get_postgres_url()
return retrieval_store.retrieve( retrieval_started = time.perf_counter()
sources = retrieval_store.retrieve(
pg_url=url, pg_url=url,
embedder=embedder, embedder=embedder,
query=query_text, query=query_text,
k=8, k=8,
expand_links=True, expand_links=True,
) )
retrieval_ms = round((time.perf_counter() - retrieval_started) * 1000, 2)
return sources, {
"embedder_get_ms": embedder_ms,
"embedder_cache_hit": cache_hit,
"retrieval_ms": retrieval_ms,
"retrieve_context_internal_ms": round(embedder_ms + retrieval_ms, 2),
}
def bootstrap_retrieval(sheet: dict, query_text: str, tool_log: list[dict]) -> None: def bootstrap_retrieval(sheet: dict, query_text: str, tool_log: list[dict]) -> None:
started_at, started_perf = _start_timing() started_at, started_perf = _start_timing()
sources = retrieve_context(query_text=query_text) sources, timing = retrieve_context(query_text=query_text)
finished_at, duration_ms = _finish_timing(started_perf) finished_at, duration_ms = _finish_timing(started_perf)
source_dump = {"sources": [source.to_string() for source in sources]} source_dump = {
"sources": [source.to_string() for source in sources],
"timing": timing,
}
append_tool_log( append_tool_log(
tool_log, tool_log,
"retrieve_context", "retrieve_context",
......
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
import logging
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from app.api import canvas, chat, health, context from app.api import canvas, chat, health, context
from app.config import get_frontend_url from app.config import get_frontend_url
from app.deterministic_services import embedding_provider
logger = logging.getLogger(__name__)
@asynccontextmanager @asynccontextmanager
async def lifespan(_app: FastAPI): async def lifespan(_app: FastAPI):
health.run_startup_checks() health.run_startup_checks()
try:
warmup_timing = embedding_provider.warmup_embedder()
logger.info("Embedding warmup finished: %s", warmup_timing)
except Exception:
logger.exception("Embedding warmup failed")
yield yield
......
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment