Commit cc608989 authored by Kantz's avatar Kantz
Browse files

retrival mit Jina hinzugefügt

parent e7f264b2
......@@ -65,7 +65,7 @@ def ingest(request: IngestRequest) -> dict:
embedder = _get_embedder()
docs = vector_store.load_docs(base_dir)
embeddings = vector_store.embed_passages(embedder, [doc.markdown for doc in docs])
embeddings = vector_store.embed_documents(embedder, [doc.markdown for doc in docs])
upserted = vector_store.upsert_docs(pg_url, docs, embeddings)
return {"status": "ok", "upserted": upserted}
......
......@@ -19,7 +19,7 @@ def get_embedding_settings() -> EmbeddingSettings:
if embedding_type == "sentence-transformer":
return EmbeddingSettings(
embedding_type=embedding_type,
model=os.getenv("SENTENCE_TRANSFORMER_MODEL", "all-MiniLM-L6-v2"),
model=os.getenv("SENTENCE_TRANSFORMER_MODEL", "tencent/KaLM-Embedding-Gemma3-12B-2511"),
target_dim=int(os.getenv("SENTENCE_TRANSFORMER_TARGET_DIM", "1024")),
)
if embedding_type == "openai-like":
......
......@@ -7,6 +7,7 @@ from enum import Enum
import httpx
from pydantic import BaseModel, Field
from sentence_transformers import SentenceTransformer
import torch
# -----------------------------
......@@ -29,7 +30,7 @@ class OpenAILikeConfig(BaseModel):
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")
target_dim: int = Field(384, description="Ziel-Dimension der Embeddings")
class EmbeddingConfig(BaseModel):
......@@ -46,7 +47,7 @@ class BaseEmbeddings:
Basisklasse für Embedding-Generierung mit gemeinsamen Methoden.
"""
def __init__(self, target_dim: int = 1024) -> None:
def __init__(self, target_dim: int = 384) -> None:
self.target_dim = target_dim
def _normalize(self, vec: List[float]) -> List[float]:
......@@ -148,18 +149,28 @@ class SentenceTransformerEmbeddings(BaseEmbeddings):
def model(self) -> SentenceTransformer:
"""Liefert das SentenceTransformer-Modell (lazy load)."""
if self._model is None:
self._model = SentenceTransformer(self.model_name)
self._model = SentenceTransformer(self.model_name, trust_remote_code=True)
self._model.max_seq_length = 512
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]
def embed_documents(self, texts: List[str]) -> List[List[float]]:
"""Generiert Embeddings für eine Liste von Texten."""
passage_embeddings = self.model.encode(
sentences=texts,
task="retrieval",
prompt_name="passage",
)
return [self._truncate([float(x) for x in emb]) for emb in passage_embeddings]
def embed_query(self, text: str) -> List[float]:
"""Generiert ein Embedding für einen einzelnen Text."""
query_embeddings = self.model.encode(
sentences=[text],
task="retrieval",
prompt_name="query",
)
return self._truncate([float(x) for x in query_embeddings[0]])
# -----------------------------
......@@ -190,4 +201,4 @@ class EmbeddingFactory:
elif config.embedding_type == EmbeddingType.SENTENCE_TRANSFORMER:
return SentenceTransformerEmbeddings(config=config)
else:
raise ValueError(f"Unsupported embedding type: {config.embedding_type}")
\ No newline at end of file
raise ValueError(f"Unsupported embedding type: {config.embedding_type}")
......@@ -108,13 +108,12 @@ def load_docs(base_dir: Path) -> List[DocRecord]:
return docs
def embed_passages(embedder: EmbeddingLike, texts: List[str]) -> List[List[float]]:
prefixed = [f"passage: {t}" for t in texts]
return embedder.embed_documents(prefixed)
def embed_documents(embedder: EmbeddingLike, texts: List[str]) -> List[List[float]]:
return embedder.embed_documents(texts)
def embed_query(embedder: EmbeddingLike, text: str) -> List[float]:
return embedder.embed_query(f"query: {text}")
return embedder.embed_query(text)
DDL = """
......@@ -135,7 +134,7 @@ CREATE TABLE IF NOT EXISTS docs (
path TEXT NOT NULL,
markdown TEXT NOT NULL,
embedding VECTOR(1024) NOT NULL
embedding VECTOR(512) NOT NULL
);
CREATE INDEX IF NOT EXISTS docs_embedding_cos_idx
......
......@@ -10,3 +10,9 @@ sympy
psycopg[binary]
pgvector
pyyaml
sentence-transformers
transformers
torch
peft
torchvision
......@@ -36,7 +36,7 @@ def cli_ingest(args: argparse.Namespace) -> None:
embedder = build_embedder()
base_dir = Path(args.base)
docs = vector_store.load_docs(base_dir)
embeddings = vector_store.embed_passages(embedder, [doc.markdown for doc in docs])
embeddings = vector_store.embed_documents(embedder, [doc.markdown for doc in docs])
pg_url = args.pg or os.getenv("POSTGRES_URL")
if not pg_url:
raise ValueError("Missing 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