Ollama embeddings
This commit is contained in:
parent
6c948a4afd
commit
4219cddea4
7 changed files with 152 additions and 1 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
0
src/haiku/rag/embeddings/__init__.py
Normal file
0
src/haiku/rag/embeddings/__init__.py
Normal file
12
src/haiku/rag/embeddings/base.py
Normal file
12
src/haiku/rag/embeddings/base.py
Normal 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."
|
||||
)
|
||||
17
src/haiku/rag/embeddings/ollama.py
Normal file
17
src/haiku/rag/embeddings/ollama.py
Normal 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
48
tests/test_embedder.py
Normal 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
74
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"
|
||||
|
|
|
|||
Loading…
Reference in a new issue