Ollama embeddings

This commit is contained in:
Yiorgis Gozadinos 2025-06-15 21:27:20 +02:00
parent 6c948a4afd
commit 4219cddea4
No known key found for this signature in database
7 changed files with 152 additions and 1 deletions

View file

@ -6,6 +6,7 @@ readme = "README.md"
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
requires-python = ">=3.13"
dependencies = [
"ollama>=0.5.1",
"pydantic>=2.11.7",
"python-dotenv>=1.1.0",
"sqlite-vec>=0.1.6",

View file

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

View file

View file

@ -0,0 +1,12 @@
import functools
class EmbedderBase:
_model: str = ""
_vector_dim: int = 0
@functools.lru_cache(maxsize=128)
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

@ -0,0 +1,17 @@
import functools
from ollama import AsyncClient
from haiku.rag.config import Config
from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase):
_model: str = Config.EMBEDDING_MODEL
_vector_dim: int = 1024
@functools.lru_cache(maxsize=128)
async def embed(self, text: str) -> list[float]:
client = AsyncClient(host=Config.OLLAMA_BASE_URL)
res = await client.embeddings(model=self._model, prompt=text)
return list(res["embedding"])

48
tests/test_embedder.py Normal file
View file

@ -0,0 +1,48 @@
import numpy as np
import pytest
from haiku.rag.embeddings.ollama import Embedder
@pytest.mark.asyncio
async def test_embedder():
embedder = Embedder()
embedding = await embedder.embed("hello world")
assert len(embedding) == embedder._vector_dim
@pytest.mark.asyncio
async def test_similarity():
embedder = Embedder()
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]
# Calculate cosine similarity
def similarities(embeddings, test_embedding):
return [
np.dot(embedding, test_embedding)
/ (np.linalg.norm(embedding) * np.linalg.norm(test_embedding))
for embedding in embeddings
]
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]

74
uv.lock
View file

@ -66,6 +66,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/78/b6/6307fbef88d9b5ee7421e68d78a9f162e0da4900bc5f5793f6d3d0e34fb8/annotated_types-0.7.0-py3-none-any.whl", hash = "sha256:1f02e8b43a8fbbc3f3e0d4f0f4bfc8131bcb4eebe8849b8e5c773f3a1c582a53", size = 13643, upload-time = "2024-05-20T21:33:24.1Z" },
]
[[package]]
name = "anyio"
version = "4.9.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "idna" },
{ name = "sniffio" },
]
sdist = { url = "https://files.pythonhosted.org/packages/95/7d/4c1bd541d4dffa1b52bd83fb8527089e097a106fc90b467a7313b105f840/anyio-4.9.0.tar.gz", hash = "sha256:673c0c244e15788651a4ff38710fea9675823028a6f08a5eda409e0c9840a028", size = 190949, upload-time = "2025-03-17T00:02:54.77Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/a1/ee/48ca1a7c89ffec8b6a0c5d02b89c305671d5ffd8d3c94acf8b8c408575bb/anyio-4.9.0-py3-none-any.whl", hash = "sha256:9f76d541cad6e36af7beb62e978876f3b41e3e04f2c1fbf0884604c0a9c4d93c", size = 100916, upload-time = "2025-03-17T00:02:52.713Z" },
]
[[package]]
name = "attrs"
version = "25.3.0"
@ -232,11 +245,21 @@ http = [
{ name = "aiohttp" },
]
[[package]]
name = "h11"
version = "0.16.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/01/ee/02a2c011bdab74c6fb3c75474d40b3052059d95df7e73351460c8588d963/h11-0.16.0.tar.gz", hash = "sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1", size = 101250, upload-time = "2025-04-24T03:35:25.427Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" },
]
[[package]]
name = "haiku-rag"
version = "0.1.0"
source = { editable = "." }
dependencies = [
{ name = "ollama" },
{ name = "pydantic" },
{ name = "python-dotenv" },
{ name = "sqlite-vec" },
@ -255,6 +278,7 @@ dev = [
[package.metadata]
requires-dist = [
{ name = "ollama", specifier = ">=0.5.1" },
{ name = "pydantic", specifier = ">=2.11.7" },
{ name = "python-dotenv", specifier = ">=1.1.0" },
{ name = "sqlite-vec", specifier = ">=0.1.6" },
@ -286,6 +310,34 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/53/bf/10ca917e335861101017ff46044c90e517b574fbb37219347b83be1952f6/hf_xet-1.1.3-cp37-abi3-win_amd64.whl", hash = "sha256:b578ae5ac9c056296bb0df9d018e597c8dc6390c5266f35b5c44696003cde9f3", size = 2310934, upload-time = "2025-06-04T00:47:29.632Z" },
]
[[package]]
name = "httpcore"
version = "1.0.9"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "certifi" },
{ name = "h11" },
]
sdist = { url = "https://files.pythonhosted.org/packages/06/94/82699a10bca87a5556c9c59b5963f2d039dbd239f25bc2a63907a05a14cb/httpcore-1.0.9.tar.gz", hash = "sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8", size = 85484, upload-time = "2025-04-24T22:06:22.219Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784, upload-time = "2025-04-24T22:06:20.566Z" },
]
[[package]]
name = "httpx"
version = "0.28.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "anyio" },
{ name = "certifi" },
{ name = "httpcore" },
{ name = "idna" },
]
sdist = { url = "https://files.pythonhosted.org/packages/b1/df/48c586a5fe32a0f01324ee087459e112ebb7224f646c0b5023f5e79e9956/httpx-0.28.1.tar.gz", hash = "sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc", size = 141406, upload-time = "2024-12-06T15:37:23.222Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" },
]
[[package]]
name = "huggingface-hub"
version = "0.33.0"
@ -430,6 +482,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/ee/e8/2c8a1c9e34d6f6d600c83d5ce5b71646c32a13f34ca5c518cc060639841c/numpy-2.3.0-cp313-cp313t-win_arm64.whl", hash = "sha256:f14e016d9409680959691c109be98c436c6249eaf7f118b424679793607b5944", size = 9935345, upload-time = "2025-06-07T14:50:02.311Z" },
]
[[package]]
name = "ollama"
version = "0.5.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "httpx" },
{ name = "pydantic" },
]
sdist = { url = "https://files.pythonhosted.org/packages/8d/96/c7fe0d2d1b3053be614822a7b722c7465161b3672ce90df71515137580a0/ollama-0.5.1.tar.gz", hash = "sha256:5a799e4dc4e7af638b11e3ae588ab17623ee019e496caaf4323efbaa8feeff93", size = 41112, upload-time = "2025-05-30T21:32:48.679Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/d6/76/3f96c8cdbf3955d7a73ee94ce3e0db0755d6de1e0098a70275940d1aff2f/ollama-0.5.1-py3-none-any.whl", hash = "sha256:4c8839f35bc173c7057b1eb2cbe7f498c1a7e134eafc9192824c8aecb3617506", size = 13369, upload-time = "2025-05-30T21:32:47.429Z" },
]
[[package]]
name = "packaging"
version = "25.0"
@ -779,6 +844,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/b7/ce/149a00dd41f10bc29e5921b496af8b574d8413afcd5e30dfa0ed46c2cc5e/six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274", size = 11050, upload-time = "2024-12-04T17:35:26.475Z" },
]
[[package]]
name = "sniffio"
version = "1.3.1"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/a2/87/a6771e1546d97e7e041b6ae58d80074f81b7d5121207425c964ddf5cfdbd/sniffio-1.3.1.tar.gz", hash = "sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc", size = 20372, upload-time = "2024-02-25T23:20:04.057Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/e9/44/75a9c9421471a6c4805dbf2356f7c181a29c1879239abab1ea2cc8f38b40/sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2", size = 10235, upload-time = "2024-02-25T23:20:01.196Z" },
]
[[package]]
name = "sqlite-vec"
version = "0.1.6"