mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
feat(embedding): add conditional dimensions parameter support for OpenAI embeddings
This commit is contained in:
parent
eff323105f
commit
11a42d02c1
5 changed files with 23 additions and 13 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ embedding_models:
|
|||
model_name: text-embedding-v4
|
||||
dimensions: 1024
|
||||
enable_cache: true
|
||||
use_dimensions: false
|
||||
|
||||
file_stores:
|
||||
default:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue