From 14816ccfaf9e234377d0ed952364213efffbe2bf Mon Sep 17 00:00:00 2001 From: "zhangyibo.114514" Date: Wed, 16 Jul 2025 13:03:27 +0800 Subject: [PATCH 1/3] =?UTF-8?q?1.=20=E8=B1=86=E5=8C=85=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E9=9D=9E=E5=A4=9A=E6=A8=A1=E6=80=81=E6=A8=A1=E5=9E=8B=202.=20e?= =?UTF-8?q?mbedder=E6=94=AF=E6=8C=81Azure=20backend?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/memos/configs/embedder.py | 4 ++++ src/memos/embedders/ark.py | 24 ++++++++++++++++++++---- src/memos/embedders/universal_api.py | 9 ++++++++- 3 files changed, 32 insertions(+), 5 deletions(-) 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, From a8adafdab2affefd7bfa082f15f74407948f2316 Mon Sep 17 00:00:00 2001 From: "zhangyibo.114514" Date: Wed, 16 Jul 2025 14:50:54 +0800 Subject: [PATCH 2/3] fix typo --- src/memos/configs/embedder.py | 2 +- src/memos/embedders/ark.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/memos/configs/embedder.py b/src/memos/configs/embedder.py index fa14d8c7f..70095a194 100644 --- a/src/memos/configs/embedder.py +++ b/src/memos/configs/embedder.py @@ -24,7 +24,7 @@ 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( + multi_modal: bool = Field( default=False, description="Whether to use multi-modal embedding (text + image) with Ark", ) diff --git a/src/memos/embedders/ark.py b/src/memos/embedders/ark.py index f74846658..544cc11ae 100644 --- a/src/memos/embedders/ark.py +++ b/src/memos/embedders/ark.py @@ -44,7 +44,7 @@ def embed(self, texts: list[str]) -> list[list[float]]: Returns: List of embeddings, each represented as a list of floats. """ - if self.config.muiti_modal: + if self.config.multi_modal: texts_input = [ MultimodalEmbeddingContentPartTextParam(text=text, type="text") for text in texts ] From f2382b8bff3ea8b5d73fa8cd8ef973f3713a3a9b Mon Sep 17 00:00:00 2001 From: "zhangyibo.114514" Date: Wed, 16 Jul 2025 15:45:45 +0800 Subject: [PATCH 3/3] add example --- examples/basic_modules/embedder.py | 21 ++++++++++++++++++++- 1 file changed, 20 insertions(+), 1 deletion(-) 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])