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
from app.deterministic_services import (
Source,
context_store,
embedding_provider,
referenz_decoder,
retrieval_store,
tool_logging,
)
from app.deterministic_services.embeddings import EmbeddingFactory
T = TypeVar("T")
......@@ -113,23 +113,40 @@ def extract_user_messages(messages: list[dict]) -> list[str]:
return user_contents
def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]:
embedder = EmbeddingFactory.create(config.get_embedding_settings())
def retrieve_context(
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()
return retrieval_store.retrieve(
retrieval_started = time.perf_counter()
sources = retrieval_store.retrieve(
pg_url=url,
embedder=embedder,
query=query_text,
k=8,
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:
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)
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(
tool_log,
"retrieve_context",
......
from contextlib import asynccontextmanager
import logging
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.api import canvas, chat, health, context
from app.config import get_frontend_url
from app.deterministic_services import embedding_provider
logger = logging.getLogger(__name__)
@asynccontextmanager
async def lifespan(_app: FastAPI):
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
......
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