Commit 459a2d24 authored by Kantz's avatar Kantz
Browse files

jina_v5

parent 01c1d3c3
...@@ -33,8 +33,8 @@ def get_embedding_settings() -> EmbeddingSettings: ...@@ -33,8 +33,8 @@ def get_embedding_settings() -> EmbeddingSettings:
return EmbeddingSettings( return EmbeddingSettings(
embedding_type=embedding_type, embedding_type=embedding_type,
model=os.getenv("SENTENCE_TRANSFORMER_MODEL", model=os.getenv("SENTENCE_TRANSFORMER_MODEL",
"jinaai/jina-embeddings-v4"), "jinaai/jina-embeddings-v5-text-small-retrieval"),
target_dim=int(os.getenv("EMBEDDING_DIM", "512")), target_dim=int(os.getenv("EMBEDDING_DIM", "1024")),
) )
if embedding_type == "openai-like": if embedding_type == "openai-like":
base_url = os.getenv("OPENAI_BASE_URL") base_url = os.getenv("OPENAI_BASE_URL")
......
...@@ -158,7 +158,9 @@ class SentenceTransformerEmbeddings(BaseEmbeddings): ...@@ -158,7 +158,9 @@ class SentenceTransformerEmbeddings(BaseEmbeddings):
"""Liefert das SentenceTransformer-Modell (lazy load).""" """Liefert das SentenceTransformer-Modell (lazy load)."""
if self._model is None: if self._model is None:
self._model = SentenceTransformer( self._model = SentenceTransformer(
self.model_name, trust_remote_code=True) self.model_name,
trust_remote_code=True,
)
self._model.max_seq_length = 512 self._model.max_seq_length = 512
return self._model return self._model
...@@ -167,7 +169,7 @@ class SentenceTransformerEmbeddings(BaseEmbeddings): ...@@ -167,7 +169,7 @@ class SentenceTransformerEmbeddings(BaseEmbeddings):
passage_embeddings = self.model.encode( passage_embeddings = self.model.encode(
sentences=texts, sentences=texts,
task="retrieval", task="retrieval",
prompt_name="passage", prompt_name="document",
) )
return [self._truncate([float(x) for x in emb]) for emb in passage_embeddings] return [self._truncate([float(x) for x in emb]) for emb in passage_embeddings]
......
...@@ -216,6 +216,10 @@ def upsert_docs(pg_url: str, docs: List[DocRecord], embeddings: List[List[float] ...@@ -216,6 +216,10 @@ def upsert_docs(pg_url: str, docs: List[DocRecord], embeddings: List[List[float]
rows = [] rows = []
for doc, emb in zip(docs, embeddings): for doc, emb in zip(docs, embeddings):
if len(emb) != embedding_dim:
raise ValueError(
f"Embedding dimension {len(emb)} does not match configured target_dim {embedding_dim}"
)
m = doc.metadata m = doc.metadata
rows.append( rows.append(
{ {
......
from __future__ import annotations
import importlib.util
import os
import sys
from pathlib import Path
import unittest
from unittest.mock import MagicMock, patch
BACKEND_ROOT = Path(__file__).resolve().parents[1]
def _load_module(name: str, relative_path: str):
path = BACKEND_ROOT / relative_path
spec = importlib.util.spec_from_file_location(name, path)
if spec is None or spec.loader is None:
raise RuntimeError(f"Failed to load module spec for {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
config = _load_module("backend_config_test_module", "app/config.py")
fake_sentence_transformers = type(sys)("sentence_transformers")
fake_sentence_transformers.SentenceTransformer = object
sys.modules.setdefault("sentence_transformers", fake_sentence_transformers)
embeddings = _load_module("backend_embeddings_test_module", "app/deterministic_services/embeddings.py")
class SentenceTransformerJinaV5Test(unittest.TestCase):
def test_config_defaults_to_jina_v5(self) -> None:
with patch.dict(os.environ, {"EMBEDDING_TYPE": "sentence-transformer"}, clear=False):
settings = config.get_embedding_settings()
self.assertEqual(settings.model, "jinaai/jina-embeddings-v5-text-small-retrieval")
self.assertEqual(settings.target_dim, 1024)
def test_embedder_uses_document_and_query_prompts(self) -> None:
fake_model = MagicMock()
fake_model.encode.side_effect = [
[[0.1, 0.2, 0.3, 0.4]],
[[0.4, 0.3, 0.2, 0.1]],
]
with patch.object(embeddings, "SentenceTransformer", return_value=fake_model) as ctor:
embedder = embeddings.SentenceTransformerEmbeddings(
embeddings.SentenceTransformerConfig(
model="jinaai/jina-embeddings-v5-text-small-retrieval",
target_dim=4,
)
)
docs = embedder.embed_documents(["doc text"])
query = embedder.embed_query("query text")
ctor.assert_called_once()
self.assertEqual(fake_model.encode.call_args_list[0].kwargs["prompt_name"], "document")
self.assertEqual(fake_model.encode.call_args_list[1].kwargs["prompt_name"], "query")
self.assertEqual(len(docs[0]), 4)
self.assertEqual(len(query), 4)
if __name__ == "__main__":
unittest.main()
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