diff --git a/mem0/configs/embeddings/base.py b/mem0/configs/embeddings/base.py index de2f4324f..c903706c6 100644 --- a/mem0/configs/embeddings/base.py +++ b/mem0/configs/embeddings/base.py @@ -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 diff --git a/mem0/configs/llms/together.py b/mem0/configs/llms/together.py new file mode 100644 index 000000000..0dbdb9c91 --- /dev/null +++ b/mem0/configs/llms/together.py @@ -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 diff --git a/mem0/embeddings/together.py b/mem0/embeddings/together.py index 830747d98..958e220a5 100644 --- a/mem0/embeddings/together.py +++ b/mem0/embeddings/together.py @@ -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): """ diff --git a/mem0/llms/together.py b/mem0/llms/together.py index 72a197bc3..67df8c824 100644 --- a/mem0/llms/together.py +++ b/mem0/llms/together.py @@ -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): """ diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 30a1079bb..ba92441c8 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -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), diff --git a/tests/embeddings/test_together_embeddings.py b/tests/embeddings/test_together_embeddings.py index ede95ec12..693a576e8 100644 --- a/tests/embeddings/test_together_embeddings.py +++ b/tests/embeddings/test_together_embeddings.py @@ -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 diff --git a/tests/llms/test_together.py b/tests/llms/test_together.py index 63428baad..b99c84cf0 100644 --- a/tests/llms/test_together.py +++ b/tests/llms/test_together.py @@ -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)