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" OLLAMA_BASE_URL: str = "http://localhost:11434"
EMBEDDING_PROVIDER: str = "ollama"
EMBEDDING_MODEL: str = "mxbai-embed-large" EMBEDDING_MODEL: str = "mxbai-embed-large"
EMBEDDING_VECTOR_DIM: int = 1024
CHUNK_SIZE: int = 256 CHUNK_SIZE: int = 256
CHUNK_OVERLAP: int = 32 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 = "" _model: str = ""
_vector_dim: int = 0 _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]: async def embed(self, text: str) -> list[float]:
raise NotImplementedError( raise NotImplementedError(
"Embedder is an abstract class. Please implement the embed method in a subclass." "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 import sqlite_vec
from haiku.rag.embeddings.ollama import Embedder from haiku.rag.embeddings import get_embedder
class Store: class Store:
@ -43,7 +43,7 @@ class Store:
""") """)
# Create vector table for chunk embeddings # Create vector table for chunk embeddings
embedder = Embedder() embedder = get_embedder()
db.execute(f""" db.execute(f"""
CREATE VIRTUAL TABLE IF NOT EXISTS chunk_embeddings USING vec0( CREATE VIRTUAL TABLE IF NOT EXISTS chunk_embeddings USING vec0(
chunk_id INTEGER PRIMARY KEY, chunk_id INTEGER PRIMARY KEY,

View file

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

View file

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