Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b8f7e1dc6e | |||
| 2ef6291307 |
@@ -20,6 +20,8 @@ class BaseEmbedderConfig(ABC):
|
||||
ollama_base_url: Optional[str] = None,
|
||||
# Openai specific
|
||||
openai_base_url: Optional[str] = None,
|
||||
# Together specific
|
||||
together_base_url: Optional[str] = None,
|
||||
# Huggingface specific
|
||||
model_kwargs: Optional[dict] = None,
|
||||
huggingface_base_url: Optional[str] = None,
|
||||
@@ -58,6 +60,8 @@ class BaseEmbedderConfig(ABC):
|
||||
:type huggingface_base_url: Optional[str], optional
|
||||
:param openai_base_url: Openai base URL to be use, defaults to "https://api.openai.com/v1"
|
||||
:type openai_base_url: Optional[str], optional
|
||||
:param together_base_url: Together base URL to be use, defaults to None
|
||||
:type together_base_url: Optional[str], optional
|
||||
:param azure_kwargs: key-value arguments for the AzureOpenAI embedding model, defaults a dict inside init
|
||||
:type azure_kwargs: Optional[Dict[str, Any]], defaults a dict inside init
|
||||
:param http_client_proxies: The proxy server settings used to create self.http_client, defaults to None
|
||||
@@ -77,6 +81,7 @@ class BaseEmbedderConfig(ABC):
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.openai_base_url = openai_base_url
|
||||
self.together_base_url = together_base_url
|
||||
self.embedding_dims = embedding_dims
|
||||
|
||||
# AzureOpenAI specific
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
from typing import Optional
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
|
||||
|
||||
class TogetherConfig(BaseLlmConfig):
|
||||
"""
|
||||
Configuration class for Together-specific parameters.
|
||||
Inherits from BaseLlmConfig and adds Together-specific settings.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
# Base parameters
|
||||
model: Optional[str] = None,
|
||||
temperature: float = 0.1,
|
||||
api_key: Optional[str] = None,
|
||||
max_tokens: int = 2000,
|
||||
top_p: float = 0.1,
|
||||
top_k: int = 1,
|
||||
enable_vision: bool = False,
|
||||
vision_details: Optional[str] = "auto",
|
||||
http_client_proxies: Optional[dict] = None,
|
||||
# Together-specific parameters
|
||||
together_base_url: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Initialize Together configuration.
|
||||
|
||||
Args:
|
||||
model: Together model to use, defaults to None
|
||||
temperature: Controls randomness, defaults to 0.1
|
||||
api_key: Together API key, defaults to None
|
||||
max_tokens: Maximum tokens to generate, defaults to 2000
|
||||
top_p: Nucleus sampling parameter, defaults to 0.1
|
||||
top_k: Top-k sampling parameter, defaults to 1
|
||||
enable_vision: Enable vision capabilities, defaults to False
|
||||
vision_details: Vision detail level, defaults to "auto"
|
||||
http_client_proxies: HTTP client proxy settings, defaults to None
|
||||
together_base_url: Together API base URL, defaults to None
|
||||
"""
|
||||
# Initialize base parameters
|
||||
super().__init__(
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
api_key=api_key,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
enable_vision=enable_vision,
|
||||
vision_details=vision_details,
|
||||
http_client_proxies=http_client_proxies,
|
||||
)
|
||||
|
||||
# Together-specific parameters
|
||||
self.together_base_url = together_base_url
|
||||
@@ -13,8 +13,9 @@ class TogetherEmbedding(EmbeddingBase):
|
||||
|
||||
self.config.model = self.config.model or "intfloat/multilingual-e5-large-instruct"
|
||||
api_key = self.config.api_key or os.getenv("TOGETHER_API_KEY")
|
||||
base_url = self.config.together_base_url or os.getenv("TOGETHER_API_BASE") or "https://api.together.ai/v1"
|
||||
self.config.embedding_dims = self.config.embedding_dims or 1024
|
||||
self.client = Together(api_key=api_key)
|
||||
self.client = Together(api_key=api_key, base_url=base_url)
|
||||
|
||||
def embed(self, text, memory_action: Optional[Literal["add", "search", "update"]] = None):
|
||||
"""
|
||||
|
||||
+24
-3
@@ -1,6 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
try:
|
||||
from together import Together
|
||||
@@ -8,19 +8,40 @@ except ImportError:
|
||||
raise ImportError("The 'together' library is required. Please install it using 'pip install together'.")
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.together import TogetherConfig
|
||||
from mem0.llms.base import LLMBase
|
||||
from mem0.memory.utils import extract_json
|
||||
|
||||
|
||||
class TogetherLLM(LLMBase):
|
||||
def __init__(self, config: Optional[BaseLlmConfig] = None):
|
||||
def __init__(self, config: Optional[Union[BaseLlmConfig, TogetherConfig, Dict]] = None):
|
||||
# Convert to TogetherConfig if needed
|
||||
if config is None:
|
||||
config = TogetherConfig()
|
||||
elif isinstance(config, dict):
|
||||
config = TogetherConfig(**config)
|
||||
elif isinstance(config, BaseLlmConfig) and not isinstance(config, TogetherConfig):
|
||||
# Convert BaseLlmConfig to TogetherConfig
|
||||
config = TogetherConfig(
|
||||
model=config.model,
|
||||
temperature=config.temperature,
|
||||
api_key=config.api_key,
|
||||
max_tokens=config.max_tokens,
|
||||
top_p=config.top_p,
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
if not self.config.model:
|
||||
self.config.model = "MiniMaxAI/MiniMax-M3"
|
||||
|
||||
api_key = self.config.api_key or os.getenv("TOGETHER_API_KEY")
|
||||
self.client = Together(api_key=api_key)
|
||||
base_url = self.config.together_base_url or os.getenv("TOGETHER_API_BASE") or "https://api.together.ai/v1"
|
||||
self.client = Together(api_key=api_key, base_url=base_url)
|
||||
|
||||
def _parse_response(self, response, tools):
|
||||
"""
|
||||
|
||||
@@ -13,6 +13,7 @@ from mem0.configs.llms.lmstudio import LMStudioConfig
|
||||
from mem0.configs.llms.minimax import MinimaxConfig
|
||||
from mem0.configs.llms.ollama import OllamaConfig
|
||||
from mem0.configs.llms.openai import OpenAIConfig
|
||||
from mem0.configs.llms.together import TogetherConfig
|
||||
from mem0.configs.llms.vllm import VllmConfig
|
||||
from mem0.configs.llms.xai import XAIConfig
|
||||
from mem0.configs.rerankers.base import BaseRerankerConfig
|
||||
@@ -43,7 +44,7 @@ class LlmFactory:
|
||||
"ollama": ("mem0.llms.ollama.OllamaLLM", OllamaConfig),
|
||||
"openai": ("mem0.llms.openai.OpenAILLM", OpenAIConfig),
|
||||
"groq": ("mem0.llms.groq.GroqLLM", BaseLlmConfig),
|
||||
"together": ("mem0.llms.together.TogetherLLM", BaseLlmConfig),
|
||||
"together": ("mem0.llms.together.TogetherLLM", TogetherConfig),
|
||||
"aws_bedrock": ("mem0.llms.aws_bedrock.AWSBedrockLLM", AWSBedrockConfig),
|
||||
"litellm": ("mem0.llms.litellm.LiteLLM", BaseLlmConfig),
|
||||
"azure_openai": ("mem0.llms.azure_openai.AzureOpenAILLM", AzureOpenAIConfig),
|
||||
|
||||
@@ -85,3 +85,29 @@ def test_explicit_config_overrides_defaults(mock_together_client):
|
||||
mock_together_client.embeddings.create.return_value = Mock(data=[Mock(embedding=[0.0] * 768)])
|
||||
embedder.embed("hello")
|
||||
mock_together_client.embeddings.create.assert_called_once_with(model="BAAI/bge-base-en-v1.5", input="hello")
|
||||
|
||||
|
||||
def test_together_embedding_base_url(monkeypatch):
|
||||
# case1: default config uses Together's official base url
|
||||
monkeypatch.delenv("TOGETHER_API_BASE", raising=False)
|
||||
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS, api_key="api_key")
|
||||
embedder = TogetherEmbedding(config)
|
||||
assert str(embedder.client.base_url) == "https://api.together.ai/v1/"
|
||||
|
||||
# case2: with env variable TOGETHER_API_BASE
|
||||
provider_base_url = "https://api.provider.com/v1/"
|
||||
monkeypatch.setenv("TOGETHER_API_BASE", provider_base_url)
|
||||
config = BaseEmbedderConfig(model=DEFAULT_MODEL, embedding_dims=DEFAULT_EMBEDDING_DIMS, api_key="api_key")
|
||||
embedder = TogetherEmbedding(config)
|
||||
assert str(embedder.client.base_url) == provider_base_url
|
||||
|
||||
# case3: with config.together_base_url (explicit config beats env var)
|
||||
config_base_url = "https://api.config.com/v1/"
|
||||
config = BaseEmbedderConfig(
|
||||
model=DEFAULT_MODEL,
|
||||
embedding_dims=DEFAULT_EMBEDDING_DIMS,
|
||||
api_key="api_key",
|
||||
together_base_url=config_base_url,
|
||||
)
|
||||
embedder = TogetherEmbedding(config)
|
||||
assert str(embedder.client.base_url) == config_base_url
|
||||
|
||||
@@ -3,6 +3,7 @@ from unittest.mock import Mock, patch
|
||||
import pytest
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.together import TogetherConfig
|
||||
from mem0.llms.together import TogetherLLM
|
||||
|
||||
|
||||
@@ -14,6 +15,38 @@ def mock_together_client():
|
||||
yield mock_client
|
||||
|
||||
|
||||
def test_together_llm_base_url(monkeypatch):
|
||||
# case1: default config uses Together's official base url
|
||||
monkeypatch.delenv("TOGETHER_API_BASE", raising=False)
|
||||
config = BaseLlmConfig(
|
||||
model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
|
||||
)
|
||||
llm = TogetherLLM(config)
|
||||
assert str(llm.client.base_url) == "https://api.together.ai/v1/"
|
||||
|
||||
# case2: with env variable TOGETHER_API_BASE
|
||||
provider_base_url = "https://api.provider.com/v1/"
|
||||
monkeypatch.setenv("TOGETHER_API_BASE", provider_base_url)
|
||||
config = TogetherConfig(
|
||||
model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
|
||||
)
|
||||
llm = TogetherLLM(config)
|
||||
assert str(llm.client.base_url) == provider_base_url
|
||||
|
||||
# case3: with config.together_base_url (explicit config beats env var)
|
||||
config_base_url = "https://api.config.com/v1/"
|
||||
config = TogetherConfig(
|
||||
model="mistralai/Mixtral-8x7B-Instruct-v0.1",
|
||||
temperature=0.7,
|
||||
max_tokens=100,
|
||||
top_p=1.0,
|
||||
api_key="api_key",
|
||||
together_base_url=config_base_url,
|
||||
)
|
||||
llm = TogetherLLM(config)
|
||||
assert str(llm.client.base_url) == config_base_url
|
||||
|
||||
|
||||
def test_generate_response_without_tools(mock_together_client):
|
||||
config = BaseLlmConfig(model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0)
|
||||
llm = TogetherLLM(config)
|
||||
|
||||
Reference in New Issue
Block a user