Introduce embedder factory

This commit is contained in:
Yiorgis Gozadinos 2025-06-18 08:26:11 +02:00
parent 1e0b6b18aa
commit 51723308ed
No known key found for this signature in database
6 changed files with 27 additions and 8 deletions

View file

@ -11,7 +11,9 @@ class AppConfig(BaseModel):
OLLAMA_BASE_URL: str = "http://localhost:11434"
EMBEDDING_PROVIDER: str = "ollama"
EMBEDDING_MODEL: str = "mxbai-embed-large"
EMBEDDING_VECTOR_DIM: int = 1024
CHUNK_SIZE: int = 256
CHUNK_OVERLAP: int = 32

View file

@ -0,0 +1,13 @@
from haiku.rag.config import Config
from haiku.rag.embeddings.base import EmbedderBase
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
def get_embedder() -> EmbedderBase:
"""
Factory function to get the appropriate embedder based on the configuration.
"""
if Config.EMBEDDING_PROVIDER == "ollama":
return OllamaEmbedder(Config.EMBEDDING_MODEL, Config.EMBEDDING_VECTOR_DIM)
raise ValueError(f"Unsupported embedding provider: {Config.EMBEDDING_PROVIDER}")

View file

@ -2,6 +2,10 @@ class EmbedderBase:
_model: str = ""
_vector_dim: int = 0
def __init__(self, model: str, vector_dim: int):
self._model = model
self._vector_dim = vector_dim
async def embed(self, text: str) -> list[float]:
raise NotImplementedError(
"Embedder is an abstract class. Please implement the embed method in a subclass."

View file

@ -5,7 +5,7 @@ from typing import Literal
import sqlite_vec
from haiku.rag.embeddings.ollama import Embedder
from haiku.rag.embeddings import get_embedder
class Store:
@ -43,7 +43,7 @@ class Store:
""")
# Create vector table for chunk embeddings
embedder = Embedder()
embedder = get_embedder()
db.execute(f"""
CREATE VIRTUAL TABLE IF NOT EXISTS chunk_embeddings USING vec0(
chunk_id INTEGER PRIMARY KEY,

View file

@ -2,7 +2,7 @@ import json
import re
from haiku.rag.chunker import chunker
from haiku.rag.embeddings.ollama import Embedder
from haiku.rag.embeddings import get_embedder
from haiku.rag.store.models.chunk import Chunk
from haiku.rag.store.repositories.base import BaseRepository
@ -10,9 +10,9 @@ from haiku.rag.store.repositories.base import BaseRepository
class ChunkRepository(BaseRepository[Chunk]):
"""Repository for Chunk database operations."""
def __init__(self, store, embedder: Embedder | None = None):
def __init__(self, store):
super().__init__(store)
self.embedder = embedder or Embedder()
self.embedder = get_embedder()
async def create(self, entity: Chunk, commit: bool = True) -> Chunk:
"""Create a chunk in the database."""

View file

@ -1,19 +1,19 @@
import numpy as np
import pytest
from haiku.rag.embeddings.ollama import Embedder
from haiku.rag.embeddings import get_embedder
@pytest.mark.asyncio
async def test_embedder():
embedder = Embedder()
embedder = get_embedder()
embedding = await embedder.embed("hello world")
assert len(embedding) == embedder._vector_dim
@pytest.mark.asyncio
async def test_similarity():
embedder = Embedder()
embedder = get_embedder()
phrases = [
"I enjoy eating great food.",
"Python is my favorite programming language.",