Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion examples/basic_modules/embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
{
Expand All @@ -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": "<YOUR_AZURE_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])
4 changes: 4 additions & 0 deletions src/memos/configs/embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
24 changes: 20 additions & 4 deletions src/memos/embedders/ark.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
9 changes: 8 additions & 1 deletion src/memos/embedders/universal_api.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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,
Expand Down