feat(embedding): add conditional dimensions parameter support for OpenAI embeddings

This commit is contained in:
jinli.yl 2026-03-02 19:27:29 +08:00
parent eff323105f
commit 11a42d02c1
5 changed files with 23 additions and 13 deletions

View file

@ -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",

View file

@ -13,6 +13,7 @@ embedding_models:
model_name: text-embedding-v4
dimensions: 1024
enable_cache: true
use_dimensions: false
file_stores:
default:

View file

@ -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

View file

@ -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:

View file

@ -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: