diff --git a/src/memos/configs/embedder.py b/src/memos/configs/embedder.py index 2d5dc61fe..fa14d8c7f 100644 --- a/src/memos/configs/embedder.py +++ b/src/memos/configs/embedder.py @@ -24,6 +24,10 @@ class ArkEmbedderConfig(BaseEmbedderConfig): default="https://ark.cn-beijing.volces.com/api/v3/", description="Base URL for Ark API" ) chunk_size: int = Field(default=1, description="Chunk size for Ark API") + muiti_modal: bool = Field( + default=False, + description="Whether to use multi-modal embedding (text + image) with Ark", + ) class SenTranEmbedderConfig(BaseEmbedderConfig): diff --git a/src/memos/embedders/ark.py b/src/memos/embedders/ark.py index cc8fba809..f74846658 100644 --- a/src/memos/embedders/ark.py +++ b/src/memos/embedders/ark.py @@ -44,10 +44,26 @@ def embed(self, texts: list[str]) -> list[list[float]]: Returns: List of embeddings, each represented as a list of floats. """ - texts_input = [ - MultimodalEmbeddingContentPartTextParam(text=text, type="text") for text in texts - ] - return self.multimodal_embeddings(texts_input, chunk_size=self.config.chunk_size) + if self.config.muiti_modal: + texts_input = [ + MultimodalEmbeddingContentPartTextParam(text=text, type="text") for text in texts + ] + return self.multimodal_embeddings(inputs=texts_input, chunk_size=self.config.chunk_size) + return self.text_embedding(texts, chunk_size=self.config.chunk_size) + + def text_embedding(self, inputs: list[str], chunk_size: int | None = None) -> list[list[float]]: + chunk_size_ = chunk_size or self.config.chunk_size + embeddings: list[list[float]] = [] + for i in range(0, len(inputs), chunk_size_): + response = self.client.embeddings.create( + model=self.config.model_name_or_path, + input=inputs[i : i + chunk_size_], + ) + + data = [response.data] if isinstance(response.data, dict) else response.data + embeddings.extend(r.embedding for r in data) + + return embeddings def multimodal_embeddings( self, inputs: list[EmbeddingInputParam], chunk_size: int | None = None diff --git a/src/memos/embedders/universal_api.py b/src/memos/embedders/universal_api.py index a2863cf5d..cd797fe26 100644 --- a/src/memos/embedders/universal_api.py +++ b/src/memos/embedders/universal_api.py @@ -1,4 +1,5 @@ from openai import OpenAI as OpenAIClient +from openai import AzureOpenAI as AzureClient from memos.configs.embedder import UniversalAPIEmbedderConfig from memos.embedders.base import BaseEmbedder @@ -11,11 +12,17 @@ def __init__(self, config: UniversalAPIEmbedderConfig): if self.provider == "openai": self.client = OpenAIClient(api_key=config.api_key, base_url=config.base_url) + elif self.provider == "azure": + self.client = AzureClient( + azure_endpoint=config.base_url, + api_version="2024-03-01-preview", + api_key=config.api_key, + ) else: raise ValueError(f"Unsupported provider: {self.provider}") def embed(self, texts: list[str]) -> list[list[float]]: - if self.provider == "openai": + if self.provider == "openai" or self.provider == "azure": response = self.client.embeddings.create( model=getattr(self.config, "model_name_or_path", "text-embedding-3-large"), input=texts,