diff --git a/examples/misc/multillm_memory.py b/examples/misc/multillm_memory.py index ff8f275e5..3389a9475 100644 --- a/examples/misc/multillm_memory.py +++ b/examples/misc/multillm_memory.py @@ -10,21 +10,20 @@ Example: GPT-4 analyzes a tech stack → Claude writes documentation → Data analyst analyzes user data → All models can reference previous research. """ -from dotenv import load_dotenv import logging -from mem0 import MemoryClient + +from dotenv import load_dotenv from litellm import completion +from mem0 import MemoryClient + load_dotenv() # Configure logging logging.basicConfig( level=logging.INFO, - format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', - handlers=[ - logging.StreamHandler(), - logging.FileHandler('research_team.log') - ] + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + handlers=[logging.StreamHandler(), logging.FileHandler("research_team.log")], ) logger = logging.getLogger(__name__) @@ -36,16 +35,16 @@ memory = MemoryClient() RESEARCH_TEAM = { "tech_analyst": { "model": "gpt-4o", - "role": "Technical Analyst - Code review, architecture, and technical decisions" + "role": "Technical Analyst - Code review, architecture, and technical decisions", }, "writer": { "model": "claude-3-5-sonnet-20241022", - "role": "Documentation Writer - Clear explanations and user guides" + "role": "Documentation Writer - Clear explanations and user guides", }, "data_analyst": { "model": "gpt-4o-mini", - "role": "Data Analyst - Insights, trends, and data-driven recommendations" - } + "role": "Data Analyst - Insights, trends, and data-driven recommendations", + }, } @@ -65,11 +64,7 @@ def get_team_knowledge(topic: str, project_id: str) -> str: return "Team Knowledge Base: Empty - starting fresh research" -def research_with_specialist( - task: str, - specialist: str, - project_id: str -) -> str: +def research_with_specialist(task: str, specialist: str, project_id: str) -> str: """Assign research task to specialist with access to team knowledge""" if specialist not in RESEARCH_TEAM: @@ -91,30 +86,20 @@ Provide actionable insights in your area of expertise.""" # Call the specialist's model response = completion( model=spec_info["model"], - messages=[ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": task} - ] + messages=[{"role": "system", "content": system_prompt}, {"role": "user", "content": task}], ) result = response.choices[0].message.content # Store research in shared knowledge base using both user_id and agent_id - research_entry = [ - {"role": "user", "content": f"Task: {task}"}, - {"role": "assistant", "content": result} - ] + research_entry = [{"role": "user", "content": f"Task: {task}"}, {"role": "assistant", "content": result}] memory.add( research_entry, user_id=project_id, # Project-level memory agent_id=specialist, # Agent-specific memory - metadata={ - "contributor": specialist, - "task_type": "research", - "model_used": spec_info["model"] - }, - output_format="v1.1" + metadata={"contributor": specialist, "task_type": "research", "model_used": spec_info["model"]}, + output_format="v1.1", ) return result @@ -155,23 +140,23 @@ def demo_research_team(): { "stage": "Technical Architecture", "specialist": "tech_analyst", - "task": "Analyze the best tech stack for a multi-tenant SaaS platform handling 10k+ users. Consider scalability, cost, and development speed." + "task": "Analyze the best tech stack for a multi-tenant SaaS platform handling 10k+ users. Consider scalability, cost, and development speed.", }, { "stage": "Product Documentation", "specialist": "writer", - "task": "Based on the technical analysis, write a clear product overview and user onboarding guide for our SaaS platform." + "task": "Based on the technical analysis, write a clear product overview and user onboarding guide for our SaaS platform.", }, { "stage": "Market Analysis", "specialist": "data_analyst", - "task": "Analyze market trends and pricing strategies for our SaaS platform. What metrics should we track?" + "task": "Analyze market trends and pricing strategies for our SaaS platform. What metrics should we track?", }, { "stage": "Strategic Decision", "specialist": "tech_analyst", - "task": "Given our technical architecture, documentation, and market analysis - what should be our MVP feature priority?" - } + "task": "Given our technical architecture, documentation, and market analysis - what should be our MVP feature priority?", + }, ] logger.info("AI Research Team: Building a SaaS Product") @@ -181,7 +166,7 @@ def demo_research_team(): logger.info(f"\nStage {i}: {step['stage']}") logger.info(f"Specialist: {step['specialist']}") - result = research_with_specialist(step['task'], step['specialist'], project) + result = research_with_specialist(step["task"], step["specialist"], project) logger.info(f"Task: {step['task']}") logger.info(f"Result: {result[:200]}...\n") diff --git a/mem0/configs/llms/anthropic.py b/mem0/configs/llms/anthropic.py new file mode 100644 index 000000000..5fd921a98 --- /dev/null +++ b/mem0/configs/llms/anthropic.py @@ -0,0 +1,56 @@ +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class AnthropicConfig(BaseLlmConfig): + """ + Configuration class for Anthropic-specific parameters. + Inherits from BaseLlmConfig and adds Anthropic-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, + # Anthropic-specific parameters + anthropic_base_url: Optional[str] = None, + ): + """ + Initialize Anthropic configuration. + + Args: + model: Anthropic model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: Anthropic 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 + anthropic_base_url: Anthropic 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, + ) + + # Anthropic-specific parameters + self.anthropic_base_url = anthropic_base_url diff --git a/mem0/configs/llms/azure.py b/mem0/configs/llms/azure.py new file mode 100644 index 000000000..f4eb859a2 --- /dev/null +++ b/mem0/configs/llms/azure.py @@ -0,0 +1,57 @@ +from typing import Any, Dict, Optional + +from mem0.configs.base import AzureConfig +from mem0.configs.llms.base import BaseLlmConfig + + +class AzureOpenAIConfig(BaseLlmConfig): + """ + Configuration class for Azure OpenAI-specific parameters. + Inherits from BaseLlmConfig and adds Azure OpenAI-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, + # Azure OpenAI-specific parameters + azure_kwargs: Optional[Dict[str, Any]] = None, + ): + """ + Initialize Azure OpenAI configuration. + + Args: + model: Azure OpenAI model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: Azure OpenAI 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 + azure_kwargs: Azure-specific configuration, 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, + ) + + # Azure OpenAI-specific parameters + self.azure_kwargs = AzureConfig(**(azure_kwargs or {})) diff --git a/mem0/configs/llms/base.py b/mem0/configs/llms/base.py index 445674bf8..55561c632 100644 --- a/mem0/configs/llms/base.py +++ b/mem0/configs/llms/base.py @@ -3,12 +3,14 @@ from typing import Dict, Optional, Union import httpx -from mem0.configs.base import AzureConfig - class BaseLlmConfig(ABC): """ - Config for LLMs. + Base configuration for LLMs with only common parameters. + Provider-specific configurations should be handled by separate config classes. + + This class contains only the parameters that are common across all LLM providers. + For provider-specific parameters, use the appropriate provider config class. """ def __init__( @@ -21,89 +23,34 @@ class BaseLlmConfig(ABC): top_k: int = 1, enable_vision: bool = False, vision_details: Optional[str] = "auto", - # Openrouter specific - models: Optional[list[str]] = None, - route: Optional[str] = "fallback", - openrouter_base_url: Optional[str] = None, - # Openai specific - openai_base_url: Optional[str] = None, - site_url: Optional[str] = None, - app_name: Optional[str] = None, - # Ollama specific - ollama_base_url: Optional[str] = None, - # AzureOpenAI specific - azure_kwargs: Optional[AzureConfig] = {}, - # AzureOpenAI specific http_client_proxies: Optional[Union[Dict, str]] = None, - # DeepSeek specific - deepseek_base_url: Optional[str] = None, - # XAI specific - xai_base_url: Optional[str] = None, - # Sarvam specific - sarvam_base_url: Optional[str] = "https://api.sarvam.ai/v1", - # LM Studio specific - lmstudio_base_url: Optional[str] = "http://localhost:1234/v1", - lmstudio_response_format: dict = None, - # vLLM specific - vllm_base_url: Optional[str] = "http://localhost:8000/v1", - # AWS Bedrock specific - aws_access_key_id: Optional[str] = None, - aws_secret_access_key: Optional[str] = None, - aws_region: Optional[str] = "us-west-2", ): """ - Initializes a configuration class instance for the LLM. + Initialize a base configuration class instance for the LLM. - :param model: Controls the OpenAI model used, defaults to None - :type model: Optional[str], optional - :param temperature: Controls the randomness of the model's output. - Higher values (closer to 1) make output more random, lower values make it more deterministic, defaults to 0 - :type temperature: float, optional - :param api_key: OpenAI API key to be use, defaults to None - :type api_key: Optional[str], optional - :param max_tokens: Controls how many tokens are generated, defaults to 2000 - :type max_tokens: int, optional - :param top_p: Controls the diversity of words. Higher values (closer to 1) make word selection more diverse, - defaults to 1 - :type top_p: float, optional - :param top_k: Controls the diversity of words. Higher values make word selection more diverse, defaults to 0 - :type top_k: int, optional - :param enable_vision: Enable vision for the LLM, defaults to False - :type enable_vision: bool, optional - :param vision_details: Details of the vision to be used [low, high, auto], defaults to "auto" - :type vision_details: Optional[str], optional - :param models: Openrouter models to use, defaults to None - :type models: Optional[list[str]], optional - :param route: Openrouter route to be used, defaults to "fallback" - :type route: Optional[str], optional - :param openrouter_base_url: Openrouter base URL to be use, defaults to "https://openrouter.ai/api/v1" - :type openrouter_base_url: Optional[str], optional - :param site_url: Openrouter site URL to use, defaults to None - :type site_url: Optional[str], optional - :param app_name: Openrouter app name to use, defaults to None - :type app_name: Optional[str], optional - :param ollama_base_url: The base URL of the LLM, defaults to None - :type ollama_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 azure_kwargs: key-value arguments for the AzureOpenAI LLM 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(s) settings used to create self.http_client, defaults to None - :type http_client_proxies: Optional[Dict | str], optional - :param deepseek_base_url: DeepSeek base URL to be use, defaults to None - :type deepseek_base_url: Optional[str], optional - :param xai_base_url: XAI base URL to be use, defaults to None - :type xai_base_url: Optional[str], optional - :param sarvam_base_url: Sarvam base URL to be use, defaults to "https://api.sarvam.ai/v1" - :type sarvam_base_url: Optional[str], optional - :param lmstudio_base_url: LM Studio base URL to be use, defaults to "http://localhost:1234/v1" - :type lmstudio_base_url: Optional[str], optional - :param lmstudio_response_format: LM Studio response format to be use, defaults to None - :type lmstudio_response_format: Optional[Dict], optional - :param vllm_base_url: vLLM base URL to be use, defaults to "http://localhost:8000/v1" - :type vllm_base_url: Optional[str], optional + Args: + model: The model identifier to use (e.g., "gpt-4o-mini", "claude-3-5-sonnet-20240620") + Defaults to None (will be set by provider-specific configs) + temperature: Controls the randomness of the model's output. + Higher values (closer to 1) make output more random, lower values make it more deterministic. + Range: 0.0 to 2.0. Defaults to 0.1 + api_key: API key for the LLM provider. If None, will try to get from environment variables. + Defaults to None + max_tokens: Maximum number of tokens to generate in the response. + Range: 1 to 4096 (varies by model). Defaults to 2000 + top_p: Nucleus sampling parameter. Controls diversity via nucleus sampling. + Higher values (closer to 1) make word selection more diverse. + Range: 0.0 to 1.0. Defaults to 0.1 + top_k: Top-k sampling parameter. Limits the number of tokens considered for each step. + Higher values make word selection more diverse. + Range: 1 to 40. Defaults to 1 + enable_vision: Whether to enable vision capabilities for the model. + Only applicable to vision-enabled models. Defaults to False + vision_details: Level of detail for vision processing. + Options: "low", "high", "auto". Defaults to "auto" + http_client_proxies: Proxy settings for HTTP client. + Can be a dict or string. Defaults to None """ - self.model = model self.temperature = temperature self.api_key = api_key @@ -112,41 +59,4 @@ class BaseLlmConfig(ABC): self.top_k = top_k self.enable_vision = enable_vision self.vision_details = vision_details - - # AzureOpenAI specific self.http_client = httpx.Client(proxies=http_client_proxies) if http_client_proxies else None - - # Openrouter specific - self.models = models - self.route = route - self.openrouter_base_url = openrouter_base_url - self.openai_base_url = openai_base_url - self.site_url = site_url - self.app_name = app_name - - # Ollama specific - self.ollama_base_url = ollama_base_url - - # DeepSeek specific - self.deepseek_base_url = deepseek_base_url - - # AzureOpenAI specific - self.azure_kwargs = AzureConfig(**azure_kwargs) or {} - - # XAI specific - self.xai_base_url = xai_base_url - - # Sarvam specific - self.sarvam_base_url = sarvam_base_url - - # LM Studio specific - self.lmstudio_base_url = lmstudio_base_url - self.lmstudio_response_format = lmstudio_response_format - - # vLLM specific - self.vllm_base_url = vllm_base_url - - # AWS Bedrock specific - self.aws_access_key_id = aws_access_key_id - self.aws_secret_access_key = aws_secret_access_key - self.aws_region = aws_region diff --git a/mem0/configs/llms/deepseek.py b/mem0/configs/llms/deepseek.py new file mode 100644 index 000000000..461b5bced --- /dev/null +++ b/mem0/configs/llms/deepseek.py @@ -0,0 +1,56 @@ +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class DeepSeekConfig(BaseLlmConfig): + """ + Configuration class for DeepSeek-specific parameters. + Inherits from BaseLlmConfig and adds DeepSeek-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, + # DeepSeek-specific parameters + deepseek_base_url: Optional[str] = None, + ): + """ + Initialize DeepSeek configuration. + + Args: + model: DeepSeek model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: DeepSeek 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 + deepseek_base_url: DeepSeek 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, + ) + + # DeepSeek-specific parameters + self.deepseek_base_url = deepseek_base_url diff --git a/mem0/configs/llms/lmstudio.py b/mem0/configs/llms/lmstudio.py new file mode 100644 index 000000000..64abdd502 --- /dev/null +++ b/mem0/configs/llms/lmstudio.py @@ -0,0 +1,59 @@ +from typing import Any, Dict, Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class LMStudioConfig(BaseLlmConfig): + """ + Configuration class for LM Studio-specific parameters. + Inherits from BaseLlmConfig and adds LM Studio-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, + # LM Studio-specific parameters + lmstudio_base_url: Optional[str] = None, + lmstudio_response_format: Optional[Dict[str, Any]] = None, + ): + """ + Initialize LM Studio configuration. + + Args: + model: LM Studio model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: LM Studio 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 + lmstudio_base_url: LM Studio base URL, defaults to None + lmstudio_response_format: LM Studio response format, 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, + ) + + # LM Studio-specific parameters + self.lmstudio_base_url = lmstudio_base_url or "http://localhost:1234/v1" + self.lmstudio_response_format = lmstudio_response_format diff --git a/mem0/configs/llms/ollama.py b/mem0/configs/llms/ollama.py new file mode 100644 index 000000000..75e1cea3f --- /dev/null +++ b/mem0/configs/llms/ollama.py @@ -0,0 +1,56 @@ +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class OllamaConfig(BaseLlmConfig): + """ + Configuration class for Ollama-specific parameters. + Inherits from BaseLlmConfig and adds Ollama-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, + # Ollama-specific parameters + ollama_base_url: Optional[str] = None, + ): + """ + Initialize Ollama configuration. + + Args: + model: Ollama model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: Ollama 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 + ollama_base_url: Ollama 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, + ) + + # Ollama-specific parameters + self.ollama_base_url = ollama_base_url diff --git a/mem0/configs/llms/openai.py b/mem0/configs/llms/openai.py new file mode 100644 index 000000000..459557e58 --- /dev/null +++ b/mem0/configs/llms/openai.py @@ -0,0 +1,71 @@ +from typing import List, Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class OpenAIConfig(BaseLlmConfig): + """ + Configuration class for OpenAI and OpenRouter-specific parameters. + Inherits from BaseLlmConfig and adds OpenAI-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, + # OpenAI-specific parameters + openai_base_url: Optional[str] = None, + models: Optional[List[str]] = None, + route: Optional[str] = "fallback", + openrouter_base_url: Optional[str] = None, + site_url: Optional[str] = None, + app_name: Optional[str] = None, + ): + """ + Initialize OpenAI configuration. + + Args: + model: OpenAI model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: OpenAI 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 + openai_base_url: OpenAI API base URL, defaults to None + models: List of models for OpenRouter, defaults to None + route: OpenRouter route strategy, defaults to "fallback" + openrouter_base_url: OpenRouter base URL, defaults to None + site_url: Site URL for OpenRouter, defaults to None + app_name: Application name for OpenRouter, 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, + ) + + # OpenAI-specific parameters + self.openai_base_url = openai_base_url + self.models = models + self.route = route + self.openrouter_base_url = openrouter_base_url + self.site_url = site_url + self.app_name = app_name diff --git a/mem0/configs/llms/vllm.py b/mem0/configs/llms/vllm.py new file mode 100644 index 000000000..45c6e2651 --- /dev/null +++ b/mem0/configs/llms/vllm.py @@ -0,0 +1,56 @@ +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class VllmConfig(BaseLlmConfig): + """ + Configuration class for vLLM-specific parameters. + Inherits from BaseLlmConfig and adds vLLM-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, + # vLLM-specific parameters + vllm_base_url: Optional[str] = None, + ): + """ + Initialize vLLM configuration. + + Args: + model: vLLM model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: vLLM 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 + vllm_base_url: vLLM 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, + ) + + # vLLM-specific parameters + self.vllm_base_url = vllm_base_url or "http://localhost:8000/v1" diff --git a/mem0/graphs/neptune/main.py b/mem0/graphs/neptune/main.py index 986668316..1cfd2e3dd 100644 --- a/mem0/graphs/neptune/main.py +++ b/mem0/graphs/neptune/main.py @@ -1,4 +1,5 @@ import logging + from .base import NeptuneBase try: diff --git a/mem0/llms/anthropic.py b/mem0/llms/anthropic.py index 5f004ae8b..61c180a0a 100644 --- a/mem0/llms/anthropic.py +++ b/mem0/llms/anthropic.py @@ -1,17 +1,37 @@ import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union try: import anthropic except ImportError: raise ImportError("The 'anthropic' library is required. Please install it using 'pip install anthropic'.") +from mem0.configs.llms.anthropic import AnthropicConfig from mem0.configs.llms.base import BaseLlmConfig from mem0.llms.base import LLMBase class AnthropicLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, AnthropicConfig, Dict]] = None): + # Convert to AnthropicConfig if needed + if config is None: + config = AnthropicConfig() + elif isinstance(config, dict): + config = AnthropicConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, AnthropicConfig): + # Convert BaseLlmConfig to AnthropicConfig + config = AnthropicConfig( + 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, + ) + super().__init__(config) if not self.config.model: @@ -26,6 +46,7 @@ class AnthropicLLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a response based on the given messages using Anthropic. @@ -35,6 +56,7 @@ class AnthropicLLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional Anthropic-specific parameters. Returns: str: The generated response. @@ -48,14 +70,16 @@ class AnthropicLLM(LLMBase): else: filtered_messages.append(message) - params = { - "model": self.config.model, - "messages": filtered_messages, - "system": system_message, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": filtered_messages, + "system": system_message, + } + ) + if tools: # TODO: Remove tools if no issues found with new memory addition logic params["tools"] = tools params["tool_choice"] = tool_choice diff --git a/mem0/llms/azure_openai.py b/mem0/llms/azure_openai.py index c736c43c5..e7802bc31 100644 --- a/mem0/llms/azure_openai.py +++ b/mem0/llms/azure_openai.py @@ -1,16 +1,36 @@ import json import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from openai import AzureOpenAI +from mem0.configs.llms.azure import AzureOpenAIConfig from mem0.configs.llms.base import BaseLlmConfig from mem0.llms.base import LLMBase from mem0.memory.utils import extract_json class AzureOpenAILLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, AzureOpenAIConfig, Dict]] = None): + # Convert to AzureOpenAIConfig if needed + if config is None: + config = AzureOpenAIConfig() + elif isinstance(config, dict): + config = AzureOpenAIConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, AzureOpenAIConfig): + # Convert BaseLlmConfig to AzureOpenAIConfig + config = AzureOpenAIConfig( + 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, + ) + super().__init__(config) # Model name should match the custom deployment name chosen for it. @@ -68,6 +88,7 @@ class AzureOpenAILLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a response based on the given messages using Azure OpenAI. @@ -77,6 +98,7 @@ class AzureOpenAILLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional Azure OpenAI-specific parameters. Returns: str: The generated response. @@ -88,23 +110,29 @@ class AzureOpenAILLM(LLMBase): messages[-1]["content"] = user_prompt - common_params = { - "model": self.config.model, - "messages": messages, - } + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) if self.config.model in {"o3-mini", "o1-preview", "o1"}: - params = common_params + # Use common params for these models + pass else: - params = { - **common_params, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } - if response_format: - params["response_format"] = response_format - if tools: # TODO: Remove tools if no issues found with new memory addition logic + # Add additional parameters for other models + params.update( + { + "temperature": self.config.temperature, + "max_tokens": self.config.max_tokens, + "top_p": self.config.top_p, + } + ) + + if tools: params["tools"] = tools params["tool_choice"] = tool_choice diff --git a/mem0/llms/base.py b/mem0/llms/base.py index 0c0a2a6e8..1cc6433a6 100644 --- a/mem0/llms/base.py +++ b/mem0/llms/base.py @@ -1,23 +1,49 @@ from abc import ABC, abstractmethod -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from mem0.configs.llms.base import BaseLlmConfig class LLMBase(ABC): - def __init__(self, config: Optional[BaseLlmConfig] = None): + """ + Base class for all LLM providers. + Handles common functionality and delegates provider-specific logic to subclasses. + """ + + def __init__(self, config: Optional[Union[BaseLlmConfig, Dict]] = None): """Initialize a base LLM class - :param config: LLM configuration option class, defaults to None - :type config: Optional[BaseLlmConfig], optional + :param config: LLM configuration option class or dict, defaults to None + :type config: Optional[Union[BaseLlmConfig, Dict]], optional """ if config is None: self.config = BaseLlmConfig() + elif isinstance(config, dict): + # Handle dict-based configuration (backward compatibility) + self.config = BaseLlmConfig(**config) else: self.config = config + # Validate configuration + self._validate_config() + + def _validate_config(self): + """ + Validate the configuration. + Override in subclasses to add provider-specific validation. + """ + if not hasattr(self.config, "model"): + raise ValueError("Configuration must have a 'model' attribute") + + if not hasattr(self.config, "api_key") and not hasattr(self.config, "api_key"): + # Check if API key is available via environment variable + # This will be handled by individual providers + pass + @abstractmethod - def generate_response(self, messages, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"): + def generate_response( + self, messages: List[Dict[str, str]], tools: Optional[List[Dict]] = None, tool_choice: str = "auto", **kwargs + ): """ Generate a response based on the given messages. @@ -25,8 +51,27 @@ class LLMBase(ABC): messages (list): List of message dicts containing 'role' and 'content'. tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional provider-specific parameters. Returns: - str: The generated response. + str or dict: The generated response. """ pass + + def _get_common_params(self, **kwargs) -> Dict: + """ + Get common parameters that most providers use. + + Returns: + Dict: Common parameters dictionary. + """ + params = { + "temperature": self.config.temperature, + "max_tokens": self.config.max_tokens, + "top_p": self.config.top_p, + } + + # Add provider-specific parameters from kwargs + params.update(kwargs) + + return params diff --git a/mem0/llms/deepseek.py b/mem0/llms/deepseek.py index 85a0417a6..8b5692f2e 100644 --- a/mem0/llms/deepseek.py +++ b/mem0/llms/deepseek.py @@ -1,16 +1,36 @@ import json import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from openai import OpenAI from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.deepseek import DeepSeekConfig from mem0.llms.base import LLMBase from mem0.memory.utils import extract_json class DeepSeekLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, DeepSeekConfig, Dict]] = None): + # Convert to DeepSeekConfig if needed + if config is None: + config = DeepSeekConfig() + elif isinstance(config, dict): + config = DeepSeekConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, DeepSeekConfig): + # Convert BaseLlmConfig to DeepSeekConfig + config = DeepSeekConfig( + 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, + ) + super().__init__(config) if not self.config.model: @@ -56,6 +76,7 @@ class DeepSeekLLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a response based on the given messages using DeepSeek. @@ -65,17 +86,19 @@ class DeepSeekLLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional DeepSeek-specific parameters. Returns: str: The generated response. """ - params = { - "model": self.config.model, - "messages": messages, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) if tools: params["tools"] = tools diff --git a/mem0/llms/lmstudio.py b/mem0/llms/lmstudio.py index cdcafe555..2c3f0a9fb 100644 --- a/mem0/llms/lmstudio.py +++ b/mem0/llms/lmstudio.py @@ -1,13 +1,35 @@ -from typing import Dict, List, Optional +import json +from typing import Dict, List, Optional, Union from openai import OpenAI from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.lmstudio import LMStudioConfig from mem0.llms.base import LLMBase +from mem0.memory.utils import extract_json class LMStudioLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, LMStudioConfig, Dict]] = None): + # Convert to LMStudioConfig if needed + if config is None: + config = LMStudioConfig() + elif isinstance(config, dict): + config = LMStudioConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, LMStudioConfig): + # Convert BaseLlmConfig to LMStudioConfig + config = LMStudioConfig( + 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, + ) + super().__init__(config) self.config.model = ( @@ -18,12 +40,43 @@ class LMStudioLLM(LLMBase): self.client = OpenAI(base_url=self.config.lmstudio_base_url, api_key=self.config.api_key) + def _parse_response(self, response, tools): + """ + Process the response based on whether tools are used or not. + + Args: + response: The raw response from API. + tools: The list of tools provided in the request. + + Returns: + str or dict: The processed response. + """ + if tools: + processed_response = { + "content": response.choices[0].message.content, + "tool_calls": [], + } + + if response.choices[0].message.tool_calls: + for tool_call in response.choices[0].message.tool_calls: + processed_response["tool_calls"].append( + { + "name": tool_call.function.name, + "arguments": json.loads(extract_json(tool_call.function.arguments)), + } + ) + + return processed_response + else: + return response.choices[0].message.content + def generate_response( self, messages: List[Dict[str, str]], - response_format: dict = {"type": "json_object"}, + response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a response based on the given messages using LM Studio. @@ -33,21 +86,32 @@ class LMStudioLLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional LM Studio-specific parameters. Returns: str: The generated response. """ - params = { - "model": self.config.model, - "messages": messages, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } - if response_format: - params["response_format"] = response_format - if self.config.lmstudio_response_format is not None: + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) + + # Handle response format - LM Studio defaults to json_object + if self.config.lmstudio_response_format: params["response_format"] = self.config.lmstudio_response_format + elif response_format: + params["response_format"] = response_format + else: + # Default to json_object for LM Studio + params["response_format"] = {"type": "json_object"} + + if tools: + params["tools"] = tools + params["tool_choice"] = tool_choice response = self.client.chat.completions.create(**params) - return response.choices[0].message.content + return self._parse_response(response, tools) diff --git a/mem0/llms/ollama.py b/mem0/llms/ollama.py index 54d8b719f..b19342143 100644 --- a/mem0/llms/ollama.py +++ b/mem0/llms/ollama.py @@ -1,4 +1,4 @@ -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union try: from ollama import Client @@ -6,25 +6,37 @@ except ImportError: raise ImportError("The 'ollama' library is required. Please install it using 'pip install ollama'.") from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.ollama import OllamaConfig from mem0.llms.base import LLMBase class OllamaLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, OllamaConfig, Dict]] = None): + # Convert to OllamaConfig if needed + if config is None: + config = OllamaConfig() + elif isinstance(config, dict): + config = OllamaConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, OllamaConfig): + # Convert BaseLlmConfig to OllamaConfig + config = OllamaConfig( + 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, + ) + super().__init__(config) if not self.config.model: self.config.model = "llama3.1:70b" - self.client = Client(host=self.config.ollama_base_url) - self._ensure_model_exists() - def _ensure_model_exists(self): - """ - Ensure the specified model exists locally. If not, pull it from Ollama. - """ - local_models = self.client.list()["models"] - if not any(model.get("name") == self.config.model for model in local_models): - self.client.pull(self.config.model) + self.client = Client(host=self.config.ollama_base_url) def _parse_response(self, response, tools): """ @@ -39,22 +51,18 @@ class OllamaLLM(LLMBase): """ if tools: processed_response = { - "content": response["message"]["content"], + "content": response["message"]["content"] if isinstance(response, dict) else response.message.content, "tool_calls": [], } - if response["message"].get("tool_calls"): - for tool_call in response["message"]["tool_calls"]: - processed_response["tool_calls"].append( - { - "name": tool_call["function"]["name"], - "arguments": tool_call["function"]["arguments"], - } - ) - + # Ollama doesn't support tool calls in the same way, so we return the content return processed_response else: - return response["message"]["content"] + # Handle both dict and object responses + if isinstance(response, dict): + return response["message"]["content"] + else: + return response.message.content def generate_response( self, @@ -62,33 +70,37 @@ class OllamaLLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ - Generate a response based on the given messages using OpenAI. + Generate a response based on the given messages using Ollama. Args: messages (list): List of message dicts containing 'role' and 'content'. response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional Ollama-specific parameters. Returns: str: The generated response. """ + # Build parameters for Ollama params = { "model": self.config.model, "messages": messages, - "options": { - "temperature": self.config.temperature, - "num_predict": self.config.max_tokens, - "top_p": self.config.top_p, - }, } - if response_format: - params["format"] = "json" - if tools: - params["tools"] = tools + # Add options for Ollama (temperature, num_predict, top_p) + options = { + "temperature": self.config.temperature, + "num_predict": self.config.max_tokens, + "top_p": self.config.top_p, + } + params["options"] = options + + # Remove OpenAI-specific parameters that Ollama doesn't support + params.pop("max_tokens", None) # Ollama uses different parameter names response = self.client.chat(**params) return self._parse_response(response, tools) diff --git a/mem0/llms/openai.py b/mem0/llms/openai.py index d94a0f3af..58d066d06 100644 --- a/mem0/llms/openai.py +++ b/mem0/llms/openai.py @@ -1,17 +1,36 @@ import json import os -import warnings -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from openai import OpenAI from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.openai import OpenAIConfig from mem0.llms.base import LLMBase from mem0.memory.utils import extract_json class OpenAILLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, OpenAIConfig, Dict]] = None): + # Convert to OpenAIConfig if needed + if config is None: + config = OpenAIConfig() + elif isinstance(config, dict): + config = OpenAIConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, OpenAIConfig): + # Convert BaseLlmConfig to OpenAIConfig + config = OpenAIConfig( + 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, + ) + super().__init__(config) if not self.config.model: @@ -26,18 +45,7 @@ class OpenAILLM(LLMBase): ) else: api_key = self.config.api_key or os.getenv("OPENAI_API_KEY") - base_url = ( - self.config.openai_base_url - or os.getenv("OPENAI_API_BASE") - or os.getenv("OPENAI_BASE_URL") - or "https://api.openai.com/v1" - ) - if os.environ.get("OPENAI_API_BASE"): - warnings.warn( - "The environment variable 'OPENAI_API_BASE' is deprecated and will be removed in the 0.1.80. " - "Please use 'OPENAI_BASE_URL' instead.", - DeprecationWarning, - ) + base_url = self.config.openai_base_url or os.getenv("OPENAI_BASE_URL") or "https://api.openai.com/v1" self.client = OpenAI(api_key=api_key, base_url=base_url) @@ -77,6 +85,7 @@ class OpenAILLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a JSON response based on the given messages using OpenAI. @@ -86,17 +95,19 @@ class OpenAILLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional OpenAI-specific parameters. Returns: json: The generated response. """ - params = { - "model": self.config.model, - "messages": messages, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) if os.getenv("OPENROUTER_API_KEY"): openrouter_params = {} diff --git a/mem0/llms/vllm.py b/mem0/llms/vllm.py index 6aa13addf..7953c4ef7 100644 --- a/mem0/llms/vllm.py +++ b/mem0/llms/vllm.py @@ -1,16 +1,36 @@ import json import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from openai import OpenAI from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.vllm import VllmConfig from mem0.llms.base import LLMBase from mem0.memory.utils import extract_json class VllmLLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, VllmConfig, Dict]] = None): + # Convert to VllmConfig if needed + if config is None: + config = VllmConfig() + elif isinstance(config, dict): + config = VllmConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, VllmConfig): + # Convert BaseLlmConfig to VllmConfig + config = VllmConfig( + 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, + ) + super().__init__(config) if not self.config.model: @@ -18,8 +38,7 @@ class VllmLLM(LLMBase): self.config.api_key = self.config.api_key or os.getenv("VLLM_API_KEY") or "vllm-api-key" base_url = self.config.vllm_base_url or os.getenv("VLLM_BASE_URL") - - self.client = OpenAI(base_url=base_url, api_key=self.config.api_key) + self.client = OpenAI(api_key=self.config.api_key, base_url=base_url) def _parse_response(self, response, tools): """ @@ -57,6 +76,7 @@ class VllmLLM(LLMBase): response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ Generate a response based on the given messages using vLLM. @@ -66,20 +86,19 @@ class VllmLLM(LLMBase): response_format (str or object, optional): Format of the response. Defaults to "text". tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional vLLM-specific parameters. Returns: str: The generated response. """ - params = { - "model": self.config.model, - "messages": messages, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } - - if response_format: - params["response_format"] = response_format + # Get common parameters + params = self._get_common_params(**kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) if tools: params["tools"] = tools diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index bc049c534..63f0dc7da 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -58,7 +58,7 @@ class MemoryGraph: # Create vector index if not exists if not any(idx.get("index_name") == "memzero" for idx in index_info["vector_index_exists"]): self.graph.query( - f"CREATE VECTOR INDEX memzero ON :Entity(embedding) WITH CONFIG {{'dimension': {embedding_dims}, 'capacity': 1000, 'metric': 'cos'}};" + f"CREATE VECTOR INDEX memzero ON :Entity(embedding) WITH CONFIG {{'dimension': {embedding_dims}, 'capacity': 1000, 'metric': 'cos'}};" ) # Create label+property index if not exists if not any( @@ -68,8 +68,7 @@ class MemoryGraph: self.graph.query("CREATE INDEX ON :Entity(user_id);") # Create label index if not exists if not any( - idx.get("index type") == "label" and idx.get("label") == "Entity" - for idx in index_info["index_exists"] + idx.get("index type") == "label" and idx.get("label") == "Entity" for idx in index_info["index_exists"] ): self.graph.query("CREATE INDEX ON :Entity;") @@ -613,7 +612,7 @@ class MemoryGraph: result = self.graph.query(cypher, params=params) return result - + def _fetch_existing_indexes(self): """ Retrieves information about existing indexes and vector indexes in the Memgraph database. @@ -621,10 +620,7 @@ class MemoryGraph: Returns: dict: A dictionary containing lists of existing indexes and vector indexes. """ - + index_exists = list(self.graph.query("SHOW INDEX INFO;")) vector_index_exists = list(self.graph.query("SHOW VECTOR INDEX INFO;")) - return { - "index_exists": index_exists, - "vector_index_exists": vector_index_exists - } + return {"index_exists": index_exists, "vector_index_exists": vector_index_exists} diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 5fe04fc6d..dd798d44d 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -1,8 +1,15 @@ import importlib -from typing import Optional +from typing import Dict, Optional, Union from mem0.configs.embeddings.base import BaseEmbedderConfig +from mem0.configs.llms.anthropic import AnthropicConfig +from mem0.configs.llms.azure import AzureOpenAIConfig from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.deepseek import DeepSeekConfig +from mem0.configs.llms.lmstudio import LMStudioConfig +from mem0.configs.llms.ollama import OllamaConfig +from mem0.configs.llms.openai import OpenAIConfig +from mem0.configs.llms.vllm import VllmConfig from mem0.embeddings.mock import MockEmbeddings @@ -13,36 +20,112 @@ def load_class(class_type): class LlmFactory: + """ + Factory for creating LLM instances with appropriate configurations. + Supports both old-style BaseLlmConfig and new provider-specific configs. + """ + + # Provider mappings with their config classes provider_to_class = { - "ollama": "mem0.llms.ollama.OllamaLLM", - "openai": "mem0.llms.openai.OpenAILLM", - "groq": "mem0.llms.groq.GroqLLM", - "together": "mem0.llms.together.TogetherLLM", - "aws_bedrock": "mem0.llms.aws_bedrock.AWSBedrockLLM", - "litellm": "mem0.llms.litellm.LiteLLM", - "azure_openai": "mem0.llms.azure_openai.AzureOpenAILLM", - "openai_structured": "mem0.llms.openai_structured.OpenAIStructuredLLM", - "anthropic": "mem0.llms.anthropic.AnthropicLLM", - "azure_openai_structured": "mem0.llms.azure_openai_structured.AzureOpenAIStructuredLLM", - "gemini": "mem0.llms.gemini.GeminiLLM", - "deepseek": "mem0.llms.deepseek.DeepSeekLLM", - "xai": "mem0.llms.xai.XAILLM", - "sarvam": "mem0.llms.sarvam.SarvamLLM", - "lmstudio": "mem0.llms.lmstudio.LMStudioLLM", - "vllm": "mem0.llms.vllm.VllmLLM", - "langchain": "mem0.llms.langchain.LangchainLLM", + "ollama": ("mem0.llms.ollama.OllamaLLM", OllamaConfig), + "openai": ("mem0.llms.openai.OpenAILLM", OpenAIConfig), + "groq": ("mem0.llms.groq.GroqLLM", BaseLlmConfig), + "together": ("mem0.llms.together.TogetherLLM", BaseLlmConfig), + "aws_bedrock": ("mem0.llms.aws_bedrock.AWSBedrockLLM", BaseLlmConfig), + "litellm": ("mem0.llms.litellm.LiteLLM", BaseLlmConfig), + "azure_openai": ("mem0.llms.azure_openai.AzureOpenAILLM", AzureOpenAIConfig), + "openai_structured": ("mem0.llms.openai_structured.OpenAIStructuredLLM", OpenAIConfig), + "anthropic": ("mem0.llms.anthropic.AnthropicLLM", AnthropicConfig), + "azure_openai_structured": ("mem0.llms.azure_openai_structured.AzureOpenAIStructuredLLM", AzureOpenAIConfig), + "gemini": ("mem0.llms.gemini.GeminiLLM", BaseLlmConfig), + "deepseek": ("mem0.llms.deepseek.DeepSeekLLM", DeepSeekConfig), + "xai": ("mem0.llms.xai.XAILLM", BaseLlmConfig), + "sarvam": ("mem0.llms.sarvam.SarvamLLM", BaseLlmConfig), + "lmstudio": ("mem0.llms.lmstudio.LMStudioLLM", LMStudioConfig), + "vllm": ("mem0.llms.vllm.VllmLLM", VllmConfig), + "langchain": ("mem0.llms.langchain.LangchainLLM", BaseLlmConfig), } @classmethod - def create(cls, provider_name, config): - class_type = cls.provider_to_class.get(provider_name) - if class_type: - llm_instance = load_class(class_type) - base_config = BaseLlmConfig(**config) - return llm_instance(base_config) - else: + def create(cls, provider_name: str, config: Optional[Union[BaseLlmConfig, Dict]] = None, **kwargs): + """ + Create an LLM instance with the appropriate configuration. + + Args: + provider_name (str): The provider name (e.g., 'openai', 'anthropic') + config: Configuration object or dict. If None, will create default config + **kwargs: Additional configuration parameters + + Returns: + Configured LLM instance + + Raises: + ValueError: If provider is not supported + """ + if provider_name not in cls.provider_to_class: raise ValueError(f"Unsupported Llm provider: {provider_name}") + class_type, config_class = cls.provider_to_class[provider_name] + llm_class = load_class(class_type) + + # Handle configuration + if config is None: + # Create default config with kwargs + config = config_class(**kwargs) + elif isinstance(config, dict): + # Merge dict config with kwargs + config.update(kwargs) + config = config_class(**config) + elif isinstance(config, BaseLlmConfig): + # Convert base config to provider-specific config if needed + if config_class != BaseLlmConfig: + # Convert to provider-specific config + config_dict = { + "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, + } + config_dict.update(kwargs) + config = config_class(**config_dict) + else: + # Use base config as-is + pass + else: + # Assume it's already the correct config type + pass + + return llm_class(config) + + @classmethod + def register_provider(cls, name: str, class_path: str, config_class=None): + """ + Register a new provider. + + Args: + name (str): Provider name + class_path (str): Full path to LLM class + config_class: Configuration class for the provider (defaults to BaseLlmConfig) + """ + if config_class is None: + config_class = BaseLlmConfig + cls.provider_to_class[name] = (class_path, config_class) + + @classmethod + def get_supported_providers(cls) -> list: + """ + Get list of supported providers. + + Returns: + list: List of supported provider names + """ + return list(cls.provider_to_class.keys()) + class EmbedderFactory: provider_to_class = { diff --git a/tests/llms/test_azure_openai.py b/tests/llms/test_azure_openai.py index 7ef86e948..dadb03c9e 100644 --- a/tests/llms/test_azure_openai.py +++ b/tests/llms/test_azure_openai.py @@ -2,7 +2,7 @@ from unittest.mock import Mock, patch import pytest -from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.azure import AzureOpenAIConfig from mem0.llms.azure_openai import AzureOpenAILLM MODEL = "gpt-4o" # or your custom deployment name @@ -20,7 +20,7 @@ def mock_openai_client(): def test_generate_response_without_tools(mock_openai_client): - config = BaseLlmConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P) + config = AzureOpenAIConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P) llm = AzureOpenAILLM(config) messages = [ {"role": "system", "content": "You are a helpful assistant."}, @@ -40,7 +40,7 @@ def test_generate_response_without_tools(mock_openai_client): def test_generate_response_with_tools(mock_openai_client): - config = BaseLlmConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P) + config = AzureOpenAIConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P) llm = AzureOpenAILLM(config) messages = [ {"role": "system", "content": "You are a helpful assistant."}, @@ -107,7 +107,7 @@ def test_generate_with_http_proxies(default_headers): patch("mem0.llms.azure_openai.AzureOpenAI") as mock_azure_openai, patch("httpx.Client", new=mock_http_client), ): - config = BaseLlmConfig( + config = AzureOpenAIConfig( model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, diff --git a/tests/llms/test_deepseek.py b/tests/llms/test_deepseek.py index b7d3f5e98..4a84079a1 100644 --- a/tests/llms/test_deepseek.py +++ b/tests/llms/test_deepseek.py @@ -4,6 +4,7 @@ from unittest.mock import Mock, patch import pytest from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.deepseek import DeepSeekConfig from mem0.llms.deepseek import DeepSeekLLM @@ -24,13 +25,13 @@ def test_deepseek_llm_base_url(): # case2: with env variable DEEPSEEK_API_BASE provider_base_url = "https://api.provider.com/v1/" os.environ["DEEPSEEK_API_BASE"] = provider_base_url - config = BaseLlmConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + config = DeepSeekConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") llm = DeepSeekLLM(config) assert str(llm.client.base_url) == provider_base_url # case3: with config.deepseek_base_url config_base_url = "https://api.config.com/v1/" - config = BaseLlmConfig( + config = DeepSeekConfig( model="deepseek-chat", temperature=0.7, max_tokens=100, diff --git a/tests/llms/test_lm_studio.py b/tests/llms/test_lm_studio.py index c2eace615..a6c956e64 100644 --- a/tests/llms/test_lm_studio.py +++ b/tests/llms/test_lm_studio.py @@ -2,7 +2,7 @@ from unittest.mock import Mock, patch import pytest -from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.lmstudio import LMStudioConfig from mem0.llms.lmstudio import LMStudioLLM @@ -18,7 +18,7 @@ def mock_lm_studio_client(): def test_generate_response_without_tools(mock_lm_studio_client): - config = BaseLlmConfig( + config = LMStudioConfig( model="lmstudio-community/Meta-Llama-3.1-8B-Instruct-GGUF/Meta-Llama-3.1-8B-Instruct-Q4_K_M.gguf", temperature=0.7, max_tokens=100, @@ -45,7 +45,7 @@ def test_generate_response_without_tools(mock_lm_studio_client): def test_generate_response_specifying_response_format(mock_lm_studio_client): - config = BaseLlmConfig( + config = LMStudioConfig( model="lmstudio-community/Meta-Llama-3.1-8B-Instruct-GGUF/Meta-Llama-3.1-8B-Instruct-Q4_K_M.gguf", temperature=0.7, max_tokens=100, diff --git a/tests/llms/test_ollama.py b/tests/llms/test_ollama.py index 0b797bfae..0f1e6ac3e 100644 --- a/tests/llms/test_ollama.py +++ b/tests/llms/test_ollama.py @@ -2,7 +2,7 @@ from unittest.mock import Mock, patch import pytest -from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.ollama import OllamaConfig from mem0.llms.ollama import OllamaLLM @@ -16,7 +16,7 @@ def mock_ollama_client(): def test_generate_response_without_tools(mock_ollama_client): - config = BaseLlmConfig(model="llama3.1:70b", temperature=0.7, max_tokens=100, top_p=1.0) + config = OllamaConfig(model="llama3.1:70b", temperature=0.7, max_tokens=100, top_p=1.0) llm = OllamaLLM(config) messages = [ {"role": "system", "content": "You are a helpful assistant."}, diff --git a/tests/llms/test_openai.py b/tests/llms/test_openai.py index cdc988646..265213b97 100644 --- a/tests/llms/test_openai.py +++ b/tests/llms/test_openai.py @@ -3,7 +3,7 @@ from unittest.mock import Mock, patch import pytest -from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.openai import OpenAIConfig from mem0.llms.openai import OpenAILLM @@ -17,22 +17,22 @@ def mock_openai_client(): def test_openai_llm_base_url(): # case1: default config: with openai official base url - config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + config = OpenAIConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") llm = OpenAILLM(config) # Note: openai client will parse the raw base_url into a URL object, which will have a trailing slash assert str(llm.client.base_url) == "https://api.openai.com/v1/" # case2: with env variable OPENAI_API_BASE provider_base_url = "https://api.provider.com/v1" - os.environ["OPENAI_API_BASE"] = provider_base_url - config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + os.environ["OPENAI_BASE_URL"] = provider_base_url + config = OpenAIConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") llm = OpenAILLM(config) # Note: openai client will parse the raw base_url into a URL object, which will have a trailing slash assert str(llm.client.base_url) == provider_base_url + "/" # case3: with config.openai_base_url config_base_url = "https://api.config.com/v1" - config = BaseLlmConfig( + config = OpenAIConfig( model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key", openai_base_url=config_base_url ) llm = OpenAILLM(config) @@ -41,7 +41,7 @@ def test_openai_llm_base_url(): def test_generate_response_without_tools(mock_openai_client): - config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0) + config = OpenAIConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0) llm = OpenAILLM(config) messages = [ {"role": "system", "content": "You are a helpful assistant."}, @@ -61,7 +61,7 @@ def test_generate_response_without_tools(mock_openai_client): def test_generate_response_with_tools(mock_openai_client): - config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0) + config = OpenAIConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0) llm = OpenAILLM(config) messages = [ {"role": "system", "content": "You are a helpful assistant."},