Commit 9af93f56 authored by Kantz's avatar Kantz
Browse files

Merge branch 'show' into 'main'

retrival_temp

See merge request kantz/tutor_react!3
parents 1ee1bde8 99d957b5
...@@ -7,6 +7,8 @@ OPENAI_CHAT_MODEL="" ...@@ -7,6 +7,8 @@ OPENAI_CHAT_MODEL=""
OPENAI_CHAT_TEMPERATURE="" OPENAI_CHAT_TEMPERATURE=""
OPENAI_EMBED_MODEL="" OPENAI_EMBED_MODEL=""
EMBEDDING_TYPE="" # "openai-like" or "sentence-transformers"
POSTGRES_URL="" POSTGRES_URL=""
OLLAMA_URL="" OLLAMA_URL=""
......
import os import os
from dataclasses import dataclass from dataclasses import dataclass
from dotenv import load_dotenv from dotenv import load_dotenv
from pydantic import BaseModel
from typing import Optional
load_dotenv() 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) @dataclass(frozen=True)
class OllamaSettings: class OllamaSettings:
...@@ -14,14 +48,6 @@ class OllamaSettings: ...@@ -14,14 +48,6 @@ class OllamaSettings:
keepalive: str | None keepalive: str | None
temperature: float | None temperature: float | None
@dataclass(frozen=True)
class EmbeddingSettings:
base_url: str
api_key: str
model: str
target_dim: int
@dataclass(frozen=True) @dataclass(frozen=True)
class OpenAIChatSettings: class OpenAIChatSettings:
base_url: str base_url: str
......
from __future__ import annotations from __future__ import annotations
import math import math
from typing import List from typing import List, Optional, Union
from enum import Enum
import httpx 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: # Konfigurationsmodelle (optional, aber empfohlen)
self.base_url = base_url.rstrip("/") # -----------------------------
self.api_key = api_key
self.model = model
self.target_dim = target_dim
self.endpoint = self._embedding_endpoint()
def _embedding_endpoint(self) -> str: class EmbeddingType(str, Enum):
if self.base_url.endswith("/embeddings"): OPENAI_LIKE = "openai-like"
return self.base_url SENTENCE_TRANSFORMER = "sentence-transformer"
if self.base_url.endswith("/v1"):
return f"{self.base_url}/embeddings"
return f"{self.base_url}/v1/embeddings" 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]: def _normalize(self, vec: List[float]) -> List[float]:
"""Normalisiert einen Vektor auf L2-Norm."""
norm = math.sqrt(sum(x * x for x in vec)) norm = math.sqrt(sum(x * x for x in vec))
if norm == 0.0: if norm == 0.0:
return vec return vec
return [x / norm for x in vec] return [x / norm for x in vec]
def _truncate(self, vec: List[float]) -> List[float]: def _truncate(self, vec: List[float]) -> List[float]:
"""Trunziert oder füllt den Vektor auf die Ziel-Dimension."""
if len(vec) < self.target_dim: if len(vec) < self.target_dim:
raise ValueError(f"Embedding dimension {len(vec)} < target {self.target_dim}") raise ValueError(f"Embedding dimension {len(vec)} < target {self.target_dim}")
if len(vec) > self.target_dim: if len(vec) > self.target_dim:
vec = vec[: self.target_dim] vec = vec[: self.target_dim]
return self._normalize(vec) 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]]: def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Ruft die externe Embedding-API auf."""
payload = { payload = {
"input": inputs, "input": inputs,
"model": self.model, "model": self.model,
...@@ -44,6 +112,7 @@ class OpenAILikeEmbeddings: ...@@ -44,6 +112,7 @@ class OpenAILikeEmbeddings:
"Content-Type": "application/json", "Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}", "Authorization": f"Bearer {self.api_key}",
} }
with httpx.Client(timeout=60.0) as client: with httpx.Client(timeout=60.0) as client:
response = client.post(self.endpoint, headers=headers, json=payload) response = client.post(self.endpoint, headers=headers, json=payload)
response.raise_for_status() response.raise_for_status()
...@@ -52,17 +121,73 @@ class OpenAILikeEmbeddings: ...@@ -52,17 +121,73 @@ class OpenAILikeEmbeddings:
if not isinstance(data, list): if not isinstance(data, list):
raise ValueError("Embedding response missing '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)) data_sorted = sorted(data, key=lambda item: item.get("index", 0))
embeddings: List[List[float]] = [] embeddings: List[List[float]] = []
for item in data_sorted: for item in data_sorted:
emb = item.get("embedding") emb = item.get("embedding")
if not isinstance(emb, list): if not isinstance(emb, list):
raise ValueError("Embedding item missing 'embedding' list.") raise ValueError("Embedding item missing 'embedding' list.")
embeddings.append(self._truncate([float(x) for x in emb])) embeddings.append(self._truncate([float(x) for x in emb]))
return embeddings return embeddings
def embed_documents(self, texts: List[str]) -> List[List[float]]:
return self._embed(texts)
def embed_query(self, text: str) -> List[float]: class SentenceTransformerEmbeddings(BaseEmbeddings):
return self._embed(text)[0] """
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 __future__ import annotations
from collections import defaultdict from typing import List
from typing import Dict, List
from app.deterministic_services import Source from app.deterministic_services import Source
from app import config 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 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]: 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() url = pg_url or config.get_postgres_url()
embedder = _get_embedder()
sources = vector_store.retrieve( sources = vector_store.retrieve(
pg_url=url, pg_url=url,
embedder=embedder, embedder=embedder,
...@@ -28,4 +20,4 @@ def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source] ...@@ -28,4 +20,4 @@ def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]
k=8, k=8,
expand_links=True, expand_links=True,
) )
return sources return sources
\ No newline at end of file
...@@ -10,22 +10,19 @@ from pathlib import Path ...@@ -10,22 +10,19 @@ from pathlib import Path
from dotenv import load_dotenv from dotenv import load_dotenv
from app import config
ROOT_DIR = Path(__file__).resolve().parents[1] ROOT_DIR = Path(__file__).resolve().parents[1]
if str(ROOT_DIR) not in sys.path: if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR)) 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 from app.deterministic_services import vector_store
def build_embedder() -> OpenAILikeEmbeddings: def build_embedder() -> BaseEmbeddings:
base_url = os.getenv("OPENAI_BASE_URL") return EmbeddingFactory.create(config.get_embedding_settings())
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 cli_init_db(args: argparse.Namespace) -> None: def cli_init_db(args: argparse.Namespace) -> None:
pg_url = args.pg or os.getenv("POSTGRES_URL") 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