From 4219cddea4e12028d95fa7e5632e33a27cc5f36a Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Sun, 15 Jun 2025 21:27:20 +0200 Subject: [PATCH] Ollama embeddings --- pyproject.toml | 1 + src/haiku/rag/config.py | 1 - src/haiku/rag/embeddings/__init__.py | 0 src/haiku/rag/embeddings/base.py | 12 +++++ src/haiku/rag/embeddings/ollama.py | 17 +++++++ tests/test_embedder.py | 48 ++++++++++++++++++ uv.lock | 74 ++++++++++++++++++++++++++++ 7 files changed, 152 insertions(+), 1 deletion(-) create mode 100644 src/haiku/rag/embeddings/__init__.py create mode 100644 src/haiku/rag/embeddings/base.py create mode 100644 src/haiku/rag/embeddings/ollama.py create mode 100644 tests/test_embedder.py diff --git a/pyproject.toml b/pyproject.toml index dfd060db..4ee7a6cb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/src/haiku/rag/config.py b/src/haiku/rag/config.py index 7eee57f3..53f0939a 100644 --- a/src/haiku/rag/config.py +++ b/src/haiku/rag/config.py @@ -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 diff --git a/src/haiku/rag/embeddings/__init__.py b/src/haiku/rag/embeddings/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/haiku/rag/embeddings/base.py b/src/haiku/rag/embeddings/base.py new file mode 100644 index 00000000..781d090f --- /dev/null +++ b/src/haiku/rag/embeddings/base.py @@ -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." + ) diff --git a/src/haiku/rag/embeddings/ollama.py b/src/haiku/rag/embeddings/ollama.py new file mode 100644 index 00000000..a9d1779a --- /dev/null +++ b/src/haiku/rag/embeddings/ollama.py @@ -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"]) diff --git a/tests/test_embedder.py b/tests/test_embedder.py new file mode 100644 index 00000000..509b0c7e --- /dev/null +++ b/tests/test_embedder.py @@ -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] diff --git a/uv.lock b/uv.lock index b87f9e21..bcbd1d96 100644 --- a/uv.lock +++ b/uv.lock @@ -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"