Simply embedders
This commit is contained in:
parent
859860318a
commit
0f8124fd98
4 changed files with 5 additions and 13 deletions
|
|
@ -1,6 +1,9 @@
|
||||||
|
from haiku.rag.config import Config
|
||||||
|
|
||||||
|
|
||||||
class EmbedderBase:
|
class EmbedderBase:
|
||||||
_model: str = ""
|
_model: str = Config.EMBEDDINGS_MODEL
|
||||||
_vector_dim: int = 0
|
_vector_dim: int = Config.EMBEDDINGS_VECTOR_DIM
|
||||||
|
|
||||||
def __init__(self, model: str, vector_dim: int):
|
def __init__(self, model: str, vector_dim: int):
|
||||||
self._model = model
|
self._model = model
|
||||||
|
|
|
||||||
|
|
@ -5,9 +5,6 @@ from haiku.rag.embeddings.base import EmbedderBase
|
||||||
|
|
||||||
|
|
||||||
class Embedder(EmbedderBase):
|
class Embedder(EmbedderBase):
|
||||||
_model: str = Config.EMBEDDINGS_MODEL
|
|
||||||
_vector_dim: int = 1024
|
|
||||||
|
|
||||||
async def embed(self, text: str) -> list[float]:
|
async def embed(self, text: str) -> list[float]:
|
||||||
client = AsyncClient(host=Config.OLLAMA_BASE_URL)
|
client = AsyncClient(host=Config.OLLAMA_BASE_URL)
|
||||||
res = await client.embeddings(model=self._model, prompt=text)
|
res = await client.embeddings(model=self._model, prompt=text)
|
||||||
|
|
|
||||||
|
|
@ -1,13 +1,9 @@
|
||||||
try:
|
try:
|
||||||
from openai import AsyncOpenAI
|
from openai import AsyncOpenAI
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.embeddings.base import EmbedderBase
|
from haiku.rag.embeddings.base import EmbedderBase
|
||||||
|
|
||||||
class Embedder(EmbedderBase):
|
class Embedder(EmbedderBase):
|
||||||
_model: str = Config.EMBEDDINGS_MODEL
|
|
||||||
_vector_dim: int = 1536
|
|
||||||
|
|
||||||
async def embed(self, text: str) -> list[float]:
|
async def embed(self, text: str) -> list[float]:
|
||||||
client = AsyncOpenAI()
|
client = AsyncOpenAI()
|
||||||
response = await client.embeddings.create(
|
response = await client.embeddings.create(
|
||||||
|
|
|
||||||
|
|
@ -1,13 +1,9 @@
|
||||||
try:
|
try:
|
||||||
from voyageai.client import Client # type: ignore
|
from voyageai.client import Client # type: ignore
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.embeddings.base import EmbedderBase
|
from haiku.rag.embeddings.base import EmbedderBase
|
||||||
|
|
||||||
class Embedder(EmbedderBase):
|
class Embedder(EmbedderBase):
|
||||||
_model: str = Config.EMBEDDINGS_MODEL
|
|
||||||
_vector_dim: int = 1024
|
|
||||||
|
|
||||||
async def embed(self, text: str) -> list[float]:
|
async def embed(self, text: str) -> list[float]:
|
||||||
client = Client()
|
client = Client()
|
||||||
res = client.embed([text], model=self._model, output_dtype="float")
|
res = client.embed([text], model=self._model, output_dtype="float")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue