Commit 99d957b5 authored by Kantz's avatar Kantz
Browse files

retrival_temp

parent 1ee1bde8
......@@ -7,6 +7,8 @@ OPENAI_CHAT_MODEL=""
OPENAI_CHAT_TEMPERATURE=""
OPENAI_EMBED_MODEL=""
EMBEDDING_TYPE="" # "openai-like" or "sentence-transformers"
POSTGRES_URL=""
OLLAMA_URL=""
......
import os
from dataclasses import dataclass
from dotenv import load_dotenv
from pydantic import BaseModel
from typing import Optional
load_dotenv()
class EmbeddingSettings(BaseModel):
embedding_type: str # "openai-like" oder "sentence-transformer"
base_url: Optional[str] = None
api_key: Optional[str] = None
model: str
target_dim: int = 1024
def get_embedding_settings() -> EmbeddingSettings:
embedding_type = os.getenv("EMBEDDING_TYPE", "openai-like")
if embedding_type == "sentence-transformer":
return EmbeddingSettings(
embedding_type=embedding_type,
model=os.getenv("SENTENCE_TRANSFORMER_MODEL", "all-MiniLM-L6-v2"),
target_dim=int(os.getenv("SENTENCE_TRANSFORMER_TARGET_DIM", "1024")),
)
if embedding_type == "openai-like":
base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY")
model = os.getenv("OPENAI_EMBED_MODEL", "text-embedding-3-large")
model_target_dim = int(os.getenv("OPENAI_EMBED_TARGET_DIM", "1024"))
if not base_url or not api_key:
raise ValueError("Missing OPENAI_BASE_URL or OPENAI_API_KEY")
return EmbeddingSettings(
embedding_type=embedding_type,
base_url=base_url,
api_key=api_key,
model=model,
target_dim=model_target_dim,
)
else:
raise ValueError(f"Unsupported EMBEDDING_TYPE: {embedding_type}")
@dataclass(frozen=True)
class OllamaSettings:
......@@ -14,14 +48,6 @@ class OllamaSettings:
keepalive: str | None
temperature: float | None
@dataclass(frozen=True)
class EmbeddingSettings:
base_url: str
api_key: str
model: str
target_dim: int
@dataclass(frozen=True)
class OpenAIChatSettings:
base_url: str
......
from __future__ import annotations
import math
from typing import List
from typing import List, Optional, Union
from enum import Enum
import httpx
from pydantic import BaseModel, Field
from sentence_transformers import SentenceTransformer
class OpenAILikeEmbeddings:
def __init__(self, base_url: str, api_key: str, model: str, target_dim: int = 1024) -> None:
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.model = model
self.target_dim = target_dim
self.endpoint = self._embedding_endpoint()
# -----------------------------
# Konfigurationsmodelle (optional, aber empfohlen)
# -----------------------------
def _embedding_endpoint(self) -> str:
if self.base_url.endswith("/embeddings"):
return self.base_url
if self.base_url.endswith("/v1"):
return f"{self.base_url}/embeddings"
return f"{self.base_url}/v1/embeddings"
class EmbeddingType(str, Enum):
OPENAI_LIKE = "openai-like"
SENTENCE_TRANSFORMER = "sentence-transformer"
class OpenAILikeConfig(BaseModel):
"""Konfiguration für OpenAI-ähnliche APIs."""
base_url: str = Field(..., description="Base URL der API (z. B. http://localhost:11434/v1)")
api_key: str = Field(..., description="API-Key (z. B. 'ollama' für Ollama)")
model: str = Field(..., description="Modellname (z. B. 'nomic-embed-text')")
target_dim: int = Field(1024, description="Ziel-Dimension der Embeddings")
class SentenceTransformerConfig(BaseModel):
"""Konfiguration für lokale SentenceTransformer-Modelle."""
model: str = Field(..., description="Name des SentenceTransformer-Modells (z. B. 'all-MiniLM-L6-v2')")
target_dim: int = Field(1024, description="Ziel-Dimension der Embeddings")
class EmbeddingConfig(BaseModel):
"""Gemeinsame Konfiguration für die Factory."""
embedding_type: EmbeddingType = Field(..., description="Typ der Embeddings")
config: Union[OpenAILikeConfig, SentenceTransformerConfig] = Field(..., description="Spezifische Konfiguration")
# -----------------------------
# Basisklasse für Embeddings
# -----------------------------
class BaseEmbeddings:
"""
Basisklasse für Embedding-Generierung mit gemeinsamen Methoden.
"""
def __init__(self, target_dim: int = 1024) -> None:
self.target_dim = target_dim
def _normalize(self, vec: List[float]) -> List[float]:
"""Normalisiert einen Vektor auf L2-Norm."""
norm = math.sqrt(sum(x * x for x in vec))
if norm == 0.0:
return vec
return [x / norm for x in vec]
def _truncate(self, vec: List[float]) -> List[float]:
"""Trunziert oder füllt den Vektor auf die Ziel-Dimension."""
if len(vec) < self.target_dim:
raise ValueError(f"Embedding dimension {len(vec)} < target {self.target_dim}")
if len(vec) > self.target_dim:
vec = vec[: self.target_dim]
return self._normalize(vec)
def embed_documents(self, texts: List[str]) -> List[List[float]]:
"""Generiert Embeddings für eine Liste von Texten."""
return self._embed(texts)
def embed_query(self, text: str) -> List[float]:
"""Generiert ein Embedding für einen einzelnen Text."""
return self._embed(text)[0]
def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Abstrakte Methode – muss in Unterklassen implementiert werden."""
raise NotImplementedError("Subclass must implement _embed method.")
# -----------------------------
# Subklassen
# -----------------------------
class OpenAILikeEmbeddings(BaseEmbeddings):
"""
Embeddings-Wrapper für OpenAI-ähnliche APIs (z. B. OpenAI, Ollama, TogetherAI).
"""
def __init__(self, config: OpenAILikeConfig) -> None:
super().__init__(target_dim=config.target_dim)
self.base_url = config.base_url.rstrip("/")
self.api_key = config.api_key
self.model = config.model
self.endpoint = self._embedding_endpoint()
def _embedding_endpoint(self) -> str:
"""Berechnet den korrekten Endpunkt für die Embedding-API."""
if self.base_url.endswith("/embeddings"):
return self.base_url
if self.base_url.endswith("/v1"):
return f"{self.base_url}/embeddings"
return f"{self.base_url}/v1/embeddings"
def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Ruft die externe Embedding-API auf."""
payload = {
"input": inputs,
"model": self.model,
......@@ -44,6 +112,7 @@ class OpenAILikeEmbeddings:
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}",
}
with httpx.Client(timeout=60.0) as client:
response = client.post(self.endpoint, headers=headers, json=payload)
response.raise_for_status()
......@@ -52,17 +121,73 @@ class OpenAILikeEmbeddings:
if not isinstance(data, list):
raise ValueError("Embedding response missing 'data' list.")
# Sortiere nach Index, falls nötig
data_sorted = sorted(data, key=lambda item: item.get("index", 0))
embeddings: List[List[float]] = []
for item in data_sorted:
emb = item.get("embedding")
if not isinstance(emb, list):
raise ValueError("Embedding item missing 'embedding' list.")
embeddings.append(self._truncate([float(x) for x in emb]))
return embeddings
def embed_documents(self, texts: List[str]) -> List[List[float]]:
return self._embed(texts)
def embed_query(self, text: str) -> List[float]:
return self._embed(text)[0]
class SentenceTransformerEmbeddings(BaseEmbeddings):
"""
Embeddings-Wrapper für lokale SentenceTransformer Modelle.
"""
def __init__(self, config: SentenceTransformerConfig) -> None:
super().__init__(target_dim=config.target_dim)
self.model_name = config.model
self._model: Optional[SentenceTransformer] = None
@property
def model(self) -> SentenceTransformer:
"""Liefert das SentenceTransformer-Modell (lazy load)."""
if self._model is None:
self._model = SentenceTransformer(self.model_name)
return self._model
def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Generiert Embeddings mit dem lokalen SentenceTransformer-Modell."""
embeddings = self.model.encode(inputs)
# Konvertiere in Liste von Listen (falls nötig)
if isinstance(embeddings, list) and all(isinstance(x, (int, float)) for x in embeddings[0]):
# Einzelner Vektor
return [self._truncate(embeddings)]
# Mehrere Vektoren
return [self._truncate(vec) for vec in embeddings]
# -----------------------------
# Factory: Erzeugt die richtige Embeddings-Instanz
# -----------------------------
class EmbeddingFactory:
"""
Factory-Klasse zur dynamischen Erzeugung von Embeddings-Instanzen.
"""
@staticmethod
def create(config: EmbeddingConfig) -> BaseEmbeddings:
"""
Erzeugt eine Embeddings-Instanz basierend auf der Konfiguration.
Args:
config (EmbeddingConfig): Die Konfiguration mit Typ und Details.
Returns:
BaseEmbeddings: Instanz der passenden Embeddings-Klasse.
Raises:
ValueError: Wenn der Typ nicht unterstützt wird.
"""
if config.embedding_type == EmbeddingType.OPENAI_LIKE:
return OpenAILikeEmbeddings(config=config.config)
elif config.embedding_type == EmbeddingType.SENTENCE_TRANSFORMER:
return SentenceTransformerEmbeddings(config=config.config)
else:
raise ValueError(f"Unsupported embedding type: {config.embedding_type}")
\ No newline at end of file
# app/deterministic_services/retrieval_service.py
from __future__ import annotations
from collections import defaultdict
from typing import Dict, List
from typing import List
from app.deterministic_services import Source
from app import config
from app.deterministic_services.embeddings import OpenAILikeEmbeddings
from app.deterministic_services.embeddings import EmbeddingFactory
from app.deterministic_services import vector_store
def _get_embedder() -> OpenAILikeEmbeddings:
settings = config.get_embedding_settings()
return OpenAILikeEmbeddings(
base_url=settings.base_url,
api_key=settings.api_key,
model=settings.model,
target_dim=settings.target_dim,
)
def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]:
embedder = EmbeddingFactory.create(config.get_embedding_config())
url = pg_url or config.get_postgres_url()
embedder = _get_embedder()
sources = vector_store.retrieve(
pg_url=url,
embedder=embedder,
......@@ -28,4 +20,4 @@ def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]
k=8,
expand_links=True,
)
return sources
return sources
\ No newline at end of file
......@@ -10,22 +10,19 @@ from pathlib import Path
from dotenv import load_dotenv
from app import config
ROOT_DIR = Path(__file__).resolve().parents[1]
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))
from app.deterministic_services.embeddings import OpenAILikeEmbeddings
from app.deterministic_services.embeddings import BaseEmbeddings, EmbeddingFactory
from app.deterministic_services import vector_store
def build_embedder() -> OpenAILikeEmbeddings:
base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY")
model = os.getenv("OPENAI_EMBED_MODEL", "text-embedding-3-large")
if not base_url or not api_key:
raise ValueError("Missing OPENAI_BASE_URL or OPENAI_API_KEY")
return OpenAILikeEmbeddings(base_url=base_url, api_key=api_key, model=model, target_dim=1024)
def build_embedder() -> BaseEmbeddings:
return EmbeddingFactory.create(config.get_embedding_settings())
def cli_init_db(args: argparse.Namespace) -> None:
pg_url = args.pg or os.getenv("POSTGRES_URL")
......
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