VLLM embedding
This commit is contained in:
parent
fcacf735e5
commit
16f7ba9d99
3 changed files with 53 additions and 0 deletions
|
|
@ -33,6 +33,9 @@ class AppConfig(BaseModel):
|
|||
CONTEXT_CHUNK_RADIUS: int = 0
|
||||
|
||||
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
||||
VLLM_EMBEDDINGS_BASE_URL: str = ""
|
||||
VLLM_RERANK_BASE_URL: str = ""
|
||||
VLLM_QA_BASE_URL: str = ""
|
||||
|
||||
# Provider keys
|
||||
VOYAGE_API_KEY: str = ""
|
||||
|
|
|
|||
16
src/haiku/rag/embeddings/vllm.py
Normal file
16
src/haiku/rag/embeddings/vllm.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
from openai import AsyncOpenAI
|
||||
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.embeddings.base import EmbedderBase
|
||||
|
||||
|
||||
class Embedder(EmbedderBase):
|
||||
async def embed(self, text: str) -> list[float]:
|
||||
client = AsyncOpenAI(
|
||||
base_url=f"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy"
|
||||
)
|
||||
response = await client.embeddings.create(
|
||||
model=self._model,
|
||||
input=text,
|
||||
)
|
||||
return response.data[0].embedding
|
||||
|
|
@ -4,9 +4,11 @@ import pytest
|
|||
from haiku.rag.config import Config
|
||||
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
|
||||
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
||||
from haiku.rag.embeddings.vllm import Embedder as VLLMEmbedder
|
||||
|
||||
OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY)
|
||||
VOYAGEAI_AVAILABLE = bool(Config.VOYAGE_API_KEY)
|
||||
VLLM_EMBEDDINGS_AVAILABLE = bool(Config.VLLM_EMBEDDINGS_BASE_URL)
|
||||
|
||||
|
||||
# Calculate cosine similarity
|
||||
|
|
@ -111,3 +113,35 @@ async def test_voyageai_embedder():
|
|||
|
||||
except ImportError:
|
||||
pytest.skip("VoyageAI package not installed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not VLLM_EMBEDDINGS_AVAILABLE, reason="vLLM embeddings server not configured"
|
||||
)
|
||||
async def test_vllm_embedder():
|
||||
embedder = VLLMEmbedder("mixedbread-ai/mxbai-embed-large-v1", 512)
|
||||
phrases = [
|
||||
"I enjoy eating great food.",
|
||||
"Python is my favorite programming language.",
|
||||
"I love to travel and see new places.",
|
||||
]
|
||||
embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases]
|
||||
|
||||
test_phrase = "I am going for a camping trip."
|
||||
test_embedding = await embedder.embed(test_phrase)
|
||||
|
||||
sims = similarities(embeddings, test_embedding)
|
||||
assert max(sims) == sims[2]
|
||||
|
||||
test_phrase = "When is dinner ready?"
|
||||
test_embedding = await embedder.embed(test_phrase)
|
||||
|
||||
sims = similarities(embeddings, test_embedding)
|
||||
assert max(sims) == sims[0]
|
||||
|
||||
test_phrase = "I work as a software developer."
|
||||
test_embedding = await embedder.embed(test_phrase)
|
||||
|
||||
sims = similarities(embeddings, test_embedding)
|
||||
assert max(sims) == sims[1]
|
||||
|
|
|
|||
Loading…
Reference in a new issue