Introduce embedder factory
This commit is contained in:
parent
1e0b6b18aa
commit
51723308ed
6 changed files with 27 additions and 8 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
Loading…
Reference in a new issue