diff --git a/reme/__init__.py b/reme/__init__.py index 5614028b..2621a539 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -8,7 +8,7 @@ from .reme import ReMe from .reme_cli import ReMeCli from .reme_fb import ReMeFb -__version__ = "0.3.0.0" +__version__ = "0.3.0.1" __all__ = [ "config", diff --git a/reme/config/file.yaml b/reme/config/file.yaml index 163fe79d..d7be65d3 100644 --- a/reme/config/file.yaml +++ b/reme/config/file.yaml @@ -13,6 +13,7 @@ embedding_models: model_name: text-embedding-v4 dimensions: 1024 enable_cache: true + use_dimensions: false file_stores: default: diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py index d29dc3fc..242734e5 100644 --- a/reme/core/embedding/base_embedding_model.py +++ b/reme/core/embedding/base_embedding_model.py @@ -30,6 +30,7 @@ class BaseEmbeddingModel(ABC): base_url: str | None = None, model_name: str = "", dimensions: int | None = 1024, + use_dimensions: bool = True, max_batch_size: int = 10, max_retries: int = 3, raise_exception: bool = True, @@ -46,6 +47,7 @@ class BaseEmbeddingModel(ABC): base_url: Base URL for the embedding service model_name: Name of the embedding model dimensions: Vector dimensions of the embeddings + use_dimensions: Whether to pass dimensions parameter to API (some APIs don't support it) max_batch_size: Maximum batch size for embedding requests max_retries: Maximum number of retry attempts on failure raise_exception: Whether to raise exceptions on failure @@ -58,6 +60,7 @@ class BaseEmbeddingModel(ABC): self._base_url: str = base_url self.model_name = model_name self.dimensions = dimensions + self.use_dimensions = use_dimensions self.max_batch_size = max_batch_size self.max_retries = max_retries self.raise_exception = raise_exception diff --git a/reme/core/embedding/openai_embedding_model.py b/reme/core/embedding/openai_embedding_model.py index efc0133d..3857f782 100644 --- a/reme/core/embedding/openai_embedding_model.py +++ b/reme/core/embedding/openai_embedding_model.py @@ -31,14 +31,17 @@ class OpenAIEmbeddingModel(BaseEmbeddingModel): async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: """Fetch embeddings from the API for a batch of strings.""" - completion = await self.client.embeddings.create( - model=self.model_name, - input=input_text, - dimensions=self.dimensions, - encoding_format=self.encoding_format, + create_kwargs: dict = { + "model": self.model_name, + "input": input_text, + "encoding_format": self.encoding_format, **self.kwargs, **kwargs, - ) + } + if self.use_dimensions: + create_kwargs["dimensions"] = self.dimensions + + completion = await self.client.embeddings.create(**create_kwargs) result_emb = [[] for _ in range(len(input_text))] for emb in completion.data: diff --git a/reme/core/embedding/openai_embedding_model_sync.py b/reme/core/embedding/openai_embedding_model_sync.py index bb4b7172..fc05de9c 100644 --- a/reme/core/embedding/openai_embedding_model_sync.py +++ b/reme/core/embedding/openai_embedding_model_sync.py @@ -14,14 +14,17 @@ class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel): def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]: """Fetch embeddings synchronously from the API for a batch of strings.""" - completion = self.client.embeddings.create( - model=self.model_name, - input=input_text, - dimensions=self.dimensions, - encoding_format=self.encoding_format, + create_kwargs: dict = { + "model": self.model_name, + "input": input_text, + "encoding_format": self.encoding_format, **self.kwargs, **kwargs, - ) + } + if self.use_dimensions: + create_kwargs["dimensions"] = self.dimensions + + completion = self.client.embeddings.create(**create_kwargs) result_emb = [[] for _ in range(len(input_text))] for emb in completion.data: