Refactored base class config for llms (#3241)
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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 {}))
|
||||
+28
-118
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -1,4 +1,5 @@
|
||||
import logging
|
||||
|
||||
from .base import NeptuneBase
|
||||
|
||||
try:
|
||||
|
||||
+30
-6
@@ -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 = {
|
||||
# Get common parameters
|
||||
params = self._get_common_params(**kwargs)
|
||||
params.update(
|
||||
{
|
||||
"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,
|
||||
}
|
||||
)
|
||||
|
||||
if tools: # TODO: Remove tools if no issues found with new memory addition logic
|
||||
params["tools"] = tools
|
||||
params["tool_choice"] = tool_choice
|
||||
|
||||
@@ -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 = {
|
||||
# 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,
|
||||
# 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 response_format:
|
||||
params["response_format"] = response_format
|
||||
if tools: # TODO: Remove tools if no issues found with new memory addition logic
|
||||
)
|
||||
|
||||
if tools:
|
||||
params["tools"] = tools
|
||||
params["tool_choice"] = tool_choice
|
||||
|
||||
|
||||
+51
-6
@@ -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
|
||||
|
||||
+29
-6
@@ -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 = {
|
||||
# Get common parameters
|
||||
params = self._get_common_params(**kwargs)
|
||||
params.update(
|
||||
{
|
||||
"model": self.config.model,
|
||||
"messages": messages,
|
||||
"temperature": self.config.temperature,
|
||||
"max_tokens": self.config.max_tokens,
|
||||
"top_p": self.config.top_p,
|
||||
}
|
||||
)
|
||||
|
||||
if tools:
|
||||
params["tools"] = tools
|
||||
|
||||
+75
-11
@@ -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 = {
|
||||
# Get common parameters
|
||||
params = self._get_common_params(**kwargs)
|
||||
params.update(
|
||||
{
|
||||
"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:
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
+40
-28
@@ -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:
|
||||
# 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": {
|
||||
}
|
||||
|
||||
# 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,
|
||||
},
|
||||
}
|
||||
if response_format:
|
||||
params["format"] = "json"
|
||||
params["options"] = options
|
||||
|
||||
if tools:
|
||||
params["tools"] = tools
|
||||
# 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)
|
||||
|
||||
+30
-19
@@ -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 = {
|
||||
# Get common parameters
|
||||
params = self._get_common_params(**kwargs)
|
||||
params.update(
|
||||
{
|
||||
"model": self.config.model,
|
||||
"messages": messages,
|
||||
"temperature": self.config.temperature,
|
||||
"max_tokens": self.config.max_tokens,
|
||||
"top_p": self.config.top_p,
|
||||
}
|
||||
)
|
||||
|
||||
if os.getenv("OPENROUTER_API_KEY"):
|
||||
openrouter_params = {}
|
||||
|
||||
+30
-11
@@ -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 = {
|
||||
# Get common parameters
|
||||
params = self._get_common_params(**kwargs)
|
||||
params.update(
|
||||
{
|
||||
"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 tools:
|
||||
params["tools"] = tools
|
||||
|
||||
@@ -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;")
|
||||
|
||||
@@ -624,7 +623,4 @@ class MemoryGraph:
|
||||
|
||||
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}
|
||||
|
||||
+108
-25
@@ -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 = {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."},
|
||||
|
||||
@@ -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."},
|
||||
|
||||
Reference in New Issue
Block a user