VoyageAI embeddings, should end up as a PR for pydantic-ai
This commit is contained in:
parent
3852e961b9
commit
86e31b8d29
2 changed files with 177 additions and 27 deletions
|
|
@ -128,14 +128,15 @@ def get_embedder(config: AppConfig = Config) -> EmbedderWrapper:
|
|||
|
||||
if provider == "voyageai":
|
||||
try:
|
||||
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
|
||||
from haiku.rag.embeddings.voyageai import VoyageAIEmbeddingModel
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"VoyageAI embedder requires the 'voyageai' package. "
|
||||
"Please install haiku.rag with the 'voyageai' extra: "
|
||||
"uv pip install haiku.rag[voyageai]"
|
||||
)
|
||||
return VoyageAIEmbedder(model_name, vector_dim, config) # type: ignore[return-value]
|
||||
model = VoyageAIEmbeddingModel(model_name)
|
||||
return EmbedderWrapper(Embedder(model), vector_dim)
|
||||
|
||||
if provider == "cohere":
|
||||
return EmbedderWrapper(Embedder(f"cohere:{model_name}"), vector_dim)
|
||||
|
|
|
|||
|
|
@ -1,33 +1,182 @@
|
|||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal, cast
|
||||
|
||||
from pydantic_ai.embeddings.base import EmbeddingModel
|
||||
from pydantic_ai.embeddings.result import EmbeddingResult, EmbedInputType
|
||||
from pydantic_ai.embeddings.settings import EmbeddingSettings
|
||||
from pydantic_ai.exceptions import ModelAPIError
|
||||
from pydantic_ai.usage import RequestUsage
|
||||
|
||||
try:
|
||||
from voyageai.client import Client # type: ignore
|
||||
from voyageai.client_async import AsyncClient
|
||||
from voyageai.error import VoyageError
|
||||
except ImportError as _import_error:
|
||||
raise ImportError(
|
||||
"Please install `voyageai` to use the VoyageAI embeddings model, "
|
||||
"you can use — `pip install voyageai`"
|
||||
) from _import_error
|
||||
|
||||
from haiku.rag.config import AppConfig
|
||||
LatestVoyageAIEmbeddingModelNames = Literal[
|
||||
"voyage-3-large",
|
||||
"voyage-3.5",
|
||||
"voyage-3.5-lite",
|
||||
"voyage-code-3",
|
||||
"voyage-finance-2",
|
||||
"voyage-law-2",
|
||||
"voyage-code-2",
|
||||
]
|
||||
"""Latest VoyageAI embedding models.
|
||||
|
||||
class Embedder:
|
||||
"""VoyageAI embedder with explicit query/document methods."""
|
||||
See [VoyageAI Embeddings](https://docs.voyageai.com/docs/embeddings)
|
||||
for available models and their capabilities.
|
||||
"""
|
||||
|
||||
def __init__(self, model: str, vector_dim: int, config: AppConfig):
|
||||
self._model = model
|
||||
self._vector_dim = vector_dim
|
||||
self._config = config
|
||||
VoyageAIEmbeddingModelName = str | LatestVoyageAIEmbeddingModelNames
|
||||
"""Possible VoyageAI embedding model names."""
|
||||
|
||||
async def embed_query(self, text: str) -> list[float]:
|
||||
"""Embed a search query."""
|
||||
client = Client()
|
||||
res = client.embed(
|
||||
[text], model=self._model, input_type="query", output_dtype="float"
|
||||
|
||||
class VoyageAIEmbeddingSettings(EmbeddingSettings, total=False):
|
||||
"""Settings used for a VoyageAI embedding model request.
|
||||
|
||||
All fields from [`EmbeddingSettings`][pydantic_ai.embeddings.EmbeddingSettings] are supported,
|
||||
plus VoyageAI-specific settings prefixed with `voyageai_`.
|
||||
"""
|
||||
|
||||
# ALL FIELDS MUST BE `voyageai_` PREFIXED SO YOU CAN MERGE THEM WITH OTHER MODELS.
|
||||
|
||||
voyageai_truncation: bool
|
||||
"""Whether to truncate inputs that exceed the model's context length.
|
||||
|
||||
Defaults to True. If False, an error is raised for inputs that are too long.
|
||||
"""
|
||||
|
||||
voyageai_output_dtype: Literal["float", "int8", "uint8", "binary", "ubinary"]
|
||||
"""The output data type for embeddings.
|
||||
|
||||
- `'float'` (default): 32-bit floats
|
||||
- `'int8'`: Signed 8-bit integers (quantized)
|
||||
- `'uint8'`: Unsigned 8-bit integers (quantized)
|
||||
- `'binary'`: Binary embeddings
|
||||
- `'ubinary'`: Unsigned binary embeddings
|
||||
"""
|
||||
|
||||
|
||||
_MAX_INPUT_TOKENS: dict[VoyageAIEmbeddingModelName, int] = {
|
||||
"voyage-3-large": 32000,
|
||||
"voyage-3.5": 32000,
|
||||
"voyage-3.5-lite": 32000,
|
||||
"voyage-code-3": 32000,
|
||||
"voyage-finance-2": 32000,
|
||||
"voyage-law-2": 16000,
|
||||
"voyage-code-2": 16000,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(init=False)
|
||||
class VoyageAIEmbeddingModel(EmbeddingModel):
|
||||
"""VoyageAI embedding model implementation.
|
||||
|
||||
VoyageAI provides state-of-the-art embedding models optimized for
|
||||
retrieval, with specialized models for code, finance, and legal domains.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from pydantic_ai.embeddings.voyageai import VoyageAIEmbeddingModel
|
||||
|
||||
model = VoyageAIEmbeddingModel('voyage-3.5')
|
||||
```
|
||||
"""
|
||||
|
||||
_model_name: VoyageAIEmbeddingModelName = field(repr=False)
|
||||
_client: AsyncClient = field(repr=False)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: VoyageAIEmbeddingModelName,
|
||||
*,
|
||||
api_key: str | None = None,
|
||||
max_retries: int = 0,
|
||||
timeout: int | None = None,
|
||||
settings: EmbeddingSettings | None = None,
|
||||
):
|
||||
"""Initialize a VoyageAI embedding model.
|
||||
|
||||
Args:
|
||||
model_name: The name of the VoyageAI model to use.
|
||||
See [VoyageAI models](https://docs.voyageai.com/docs/embeddings)
|
||||
for available options.
|
||||
api_key: The VoyageAI API key. If not provided, uses the
|
||||
`VOYAGE_API_KEY` environment variable.
|
||||
max_retries: Maximum number of retries for failed requests.
|
||||
timeout: Request timeout in seconds.
|
||||
settings: Model-specific [`EmbeddingSettings`][pydantic_ai.embeddings.EmbeddingSettings]
|
||||
to use as defaults for this model.
|
||||
"""
|
||||
self._model_name = model_name
|
||||
self._client = AsyncClient(
|
||||
api_key=api_key,
|
||||
max_retries=max_retries,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
super().__init__(settings=settings)
|
||||
|
||||
@property
|
||||
def model_name(self) -> VoyageAIEmbeddingModelName:
|
||||
"""The embedding model name."""
|
||||
return self._model_name
|
||||
|
||||
@property
|
||||
def system(self) -> str:
|
||||
"""The embedding model provider."""
|
||||
return "voyageai"
|
||||
|
||||
async def embed(
|
||||
self,
|
||||
inputs: str | Sequence[str],
|
||||
*,
|
||||
input_type: EmbedInputType,
|
||||
settings: EmbeddingSettings | None = None,
|
||||
) -> EmbeddingResult:
|
||||
inputs, settings = self.prepare_embed(inputs, settings)
|
||||
settings = cast(VoyageAIEmbeddingSettings, settings)
|
||||
|
||||
voyageai_input_type = "document" if input_type == "document" else "query"
|
||||
|
||||
try:
|
||||
response = await self._client.embed(
|
||||
texts=list(inputs),
|
||||
model=self.model_name,
|
||||
input_type=voyageai_input_type,
|
||||
truncation=settings.get("voyageai_truncation", True),
|
||||
output_dtype=settings.get("voyageai_output_dtype", "float"),
|
||||
output_dimension=settings.get("dimensions"),
|
||||
)
|
||||
return res.embeddings[0] # type: ignore[return-value]
|
||||
except VoyageError as e:
|
||||
raise ModelAPIError(model_name=self.model_name, message=str(e)) from e
|
||||
|
||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||
"""Embed documents/chunks for indexing."""
|
||||
if not texts:
|
||||
return []
|
||||
client = Client()
|
||||
res = client.embed(
|
||||
texts, model=self._model, input_type="document", output_dtype="float"
|
||||
)
|
||||
return res.embeddings # type: ignore[return-value]
|
||||
return EmbeddingResult(
|
||||
embeddings=response.embeddings,
|
||||
inputs=inputs,
|
||||
input_type=input_type,
|
||||
usage=_map_usage(response.total_tokens, self.model_name),
|
||||
model_name=self.model_name,
|
||||
provider_name=self.system,
|
||||
)
|
||||
|
||||
except ImportError:
|
||||
pass
|
||||
async def max_input_tokens(self) -> int | None:
|
||||
return _MAX_INPUT_TOKENS.get(self.model_name)
|
||||
|
||||
|
||||
def _map_usage(total_tokens: int, model: str) -> RequestUsage:
|
||||
usage_data = {"total_tokens": total_tokens}
|
||||
response_data = {"model": model, "usage": usage_data}
|
||||
|
||||
return RequestUsage.extract(
|
||||
response_data,
|
||||
provider="voyageai",
|
||||
provider_url="https://api.voyageai.com",
|
||||
provider_fallback="voyageai",
|
||||
api_flavor="embeddings",
|
||||
)
|
||||
|
|
|
|||
Loading…
Reference in a new issue