Skip to content
Closed
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
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")
muiti_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.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
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