diff --git a/examples/basic_modules/embedder.py b/examples/basic_modules/embedder.py index f35746415..7cc7942da 100644 --- a/examples/basic_modules/embedder.py +++ b/examples/basic_modules/embedder.py @@ -48,7 +48,7 @@ print("Scenario 3 HF embedding shape:", len(embedding_hf[0])) print("==" * 20) -# === Scenario 4: Using UniversalAPIEmbedder === +# === Scenario 4: Using UniversalAPIEmbedder(OpenAI) === config_api = EmbedderConfigFactory.model_validate( { @@ -66,3 +66,22 @@ embedding_api = embedder_api.embed([text_api]) print("Scenario 4: OpenAI API embedding vector length:", len(embedding_api[0])) print("Embedding preview:", embedding_api[0][:10]) + +# === Scenario 5: Using UniversalAPIEmbedder(Azure) === + +config_api = EmbedderConfigFactory.model_validate( + { + "backend": "universal_api", + "config": { + "provider": "azure", + "api_key": "", + "model_name_or_path": "text-embedding-3-large", + "base_url": "https://open.azure.com/openapi/online/v2/", + }, + } +) +embedder_api = EmbedderFactory.from_config(config_api) +text_api = "This is a sample text for embedding generation using Azure API." +embedding_api = embedder_api.embed([text_api]) +print("Scenario 5: Azure API embedding vector length:", len(embedding_api[0])) +print("Embedding preview:", embedding_api[0][:10]) diff --git a/src/memos/configs/embedder.py b/src/memos/configs/embedder.py index 2d5dc61fe..70095a194 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") + multi_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..544cc11ae 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.multi_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,