fix(bedrock): omit topP for Anthropic Converse; use AWSBedrockConfig in LlmFactory (#4469)
This commit is contained in:
@@ -16,7 +16,7 @@ class AWSBedrockConfig(BaseLlmConfig):
|
||||
model: Optional[str] = None,
|
||||
temperature: float = 0.1,
|
||||
max_tokens: int = 2000,
|
||||
top_p: float = 0.9,
|
||||
top_p: Optional[float] = None,
|
||||
top_k: int = 1,
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
@@ -33,7 +33,8 @@ class AWSBedrockConfig(BaseLlmConfig):
|
||||
model: Bedrock model identifier (e.g., "amazon.nova-3-mini-20241119-v1:0")
|
||||
temperature: Controls randomness (0.0 to 2.0)
|
||||
max_tokens: Maximum tokens to generate
|
||||
top_p: Nucleus sampling parameter (0.0 to 1.0)
|
||||
top_p: Nucleus sampling (0.0–1.0). Default None omits topP on Converse
|
||||
(required for Anthropic, which rejects temperature and topP together).
|
||||
top_k: Top-k sampling parameter (1 to 40)
|
||||
aws_access_key_id: AWS access key (optional, uses env vars if not provided)
|
||||
aws_secret_access_key: AWS secret key (optional, uses env vars if not provided)
|
||||
@@ -75,13 +76,16 @@ class AWSBedrockConfig(BaseLlmConfig):
|
||||
|
||||
def get_model_config(self) -> Dict[str, Any]:
|
||||
"""Get model-specific configuration parameters."""
|
||||
base_config = {
|
||||
base_config: Dict[str, Any] = {
|
||||
"temperature": self.temperature,
|
||||
"max_tokens": self.max_tokens,
|
||||
"top_p": self.top_p,
|
||||
"top_k": self.top_k,
|
||||
}
|
||||
|
||||
# Only include top_p when explicitly set by the user.
|
||||
if self.top_p is not None:
|
||||
base_config["top_p"] = self.top_p
|
||||
|
||||
# Add custom model kwargs
|
||||
base_config.update(self.model_kwargs)
|
||||
|
||||
|
||||
+41
-39
@@ -228,6 +228,12 @@ class AWSBedrockLLM(LLMBase):
|
||||
|
||||
return "\n\nHuman: " + "".join(formatted_messages) + "\n\nAssistant:"
|
||||
|
||||
def _merge_optional_top_p(self, target: Dict[str, Any], *, key: str = "top_p") -> None:
|
||||
"""Add nucleus sampling to ``target`` only when ``model_config`` has ``top_p`` set."""
|
||||
top_p = self.model_config.get("top_p")
|
||||
if top_p is not None:
|
||||
target[key] = top_p
|
||||
|
||||
def _prepare_input(self, prompt: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Prepare input for the current provider's model.
|
||||
@@ -268,44 +274,38 @@ class AWSBedrockLLM(LLMBase):
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
self._merge_optional_top_p(input_body)
|
||||
else:
|
||||
# Legacy Amazon models
|
||||
input_body = {
|
||||
"inputText": prompt,
|
||||
"textGenerationConfig": {
|
||||
text_gen_config: Dict[str, Any] = {
|
||||
"maxTokenCount": self.model_config.get("max_tokens", 5000),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
},
|
||||
}
|
||||
# Remove None values
|
||||
input_body["textGenerationConfig"] = {
|
||||
k: v for k, v in input_body["textGenerationConfig"].items() if v is not None
|
||||
}
|
||||
self._merge_optional_top_p(text_gen_config, key="topP")
|
||||
input_body = {"inputText": prompt, "textGenerationConfig": text_gen_config}
|
||||
elif self.provider == "anthropic":
|
||||
input_body = {
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": prompt}]}],
|
||||
"max_tokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
}
|
||||
self._merge_optional_top_p(input_body)
|
||||
elif self.provider == "meta":
|
||||
input_body = {
|
||||
"prompt": prompt,
|
||||
"max_gen_len": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
self._merge_optional_top_p(input_body)
|
||||
elif self.provider == "mistral":
|
||||
input_body = {
|
||||
"prompt": prompt,
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
self._merge_optional_top_p(input_body)
|
||||
else:
|
||||
# Generic case - add all model config parameters
|
||||
input_body.update(self.model_config)
|
||||
@@ -479,6 +479,29 @@ class AWSBedrockLLM(LLMBase):
|
||||
|
||||
return converse_tools
|
||||
|
||||
def _default_max_tokens_for_converse(self) -> int:
|
||||
"""Default maxTokens if ``max_tokens`` is missing (Nova: 5000, else 2000)."""
|
||||
model_id = (self.config.model or "").lower()
|
||||
if self.provider == "amazon" and "nova" in model_id:
|
||||
return 5000
|
||||
return 2000
|
||||
|
||||
def _build_inference_config(self) -> Dict[str, Any]:
|
||||
"""Build Converse ``inferenceConfig``. Anthropic allows only one of temperature or topP; we keep temperature and omit topP."""
|
||||
inference_config: Dict[str, Any] = {
|
||||
"maxTokens": self.model_config.get("max_tokens", self._default_max_tokens_for_converse()),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
}
|
||||
|
||||
top_p = self.model_config.get("top_p")
|
||||
if top_p is not None:
|
||||
if self.provider == "anthropic":
|
||||
logger.debug("Omitting topP for Anthropic Converse (using temperature); top_p=%s", top_p)
|
||||
else:
|
||||
inference_config["topP"] = top_p
|
||||
|
||||
return inference_config
|
||||
|
||||
def _generate_with_tools(self, messages: List[Dict[str, str]], tools: List[Dict], stream: bool = False) -> Dict[str, Any]:
|
||||
"""Generate response with tool calling support using correct message format."""
|
||||
# Format messages for tool-enabled models
|
||||
@@ -501,11 +524,7 @@ class AWSBedrockLLM(LLMBase):
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
"inferenceConfig": self._build_inference_config(),
|
||||
}
|
||||
|
||||
# Add system message if present (for Anthropic)
|
||||
@@ -531,11 +550,7 @@ class AWSBedrockLLM(LLMBase):
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
"inferenceConfig": self._build_inference_config(),
|
||||
}
|
||||
|
||||
# Add system message if present
|
||||
@@ -554,26 +569,13 @@ class AWSBedrockLLM(LLMBase):
|
||||
return str(response)
|
||||
|
||||
elif self.provider == "amazon" and "nova" in self.config.model.lower():
|
||||
# Nova models use converse API even without tools
|
||||
# Nova models use the Converse API even without tools
|
||||
formatted_messages = self._format_messages_amazon(messages)
|
||||
input_body = {
|
||||
"messages": formatted_messages,
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
|
||||
# Use converse API for Nova models
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=input_body["messages"],
|
||||
inferenceConfig={
|
||||
"maxTokens": input_body["max_tokens"],
|
||||
"temperature": input_body["temperature"],
|
||||
"topP": input_body["top_p"],
|
||||
}
|
||||
messages=formatted_messages,
|
||||
inferenceConfig=self._build_inference_config(),
|
||||
)
|
||||
|
||||
return self._parse_response(response)
|
||||
else:
|
||||
# For other providers and legacy Amazon models (like Titan)
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Dict, Optional, Union
|
||||
|
||||
from mem0.configs.embeddings.base import BaseEmbedderConfig
|
||||
from mem0.configs.llms.anthropic import AnthropicConfig
|
||||
from mem0.configs.llms.aws_bedrock import AWSBedrockConfig
|
||||
from mem0.configs.llms.azure import AzureOpenAIConfig
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.deepseek import DeepSeekConfig
|
||||
@@ -38,7 +39,7 @@ class LlmFactory:
|
||||
"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),
|
||||
"aws_bedrock": ("mem0.llms.aws_bedrock.AWSBedrockLLM", AWSBedrockConfig),
|
||||
"litellm": ("mem0.llms.litellm.LiteLLM", BaseLlmConfig),
|
||||
"azure_openai": ("mem0.llms.azure_openai.AzureOpenAILLM", AzureOpenAIConfig),
|
||||
"openai_structured": ("mem0.llms.openai_structured.OpenAIStructuredLLM", OpenAIConfig),
|
||||
|
||||
@@ -0,0 +1,330 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.configs.llms.aws_bedrock import AWSBedrockConfig
|
||||
from mem0.llms.aws_bedrock import AWSBedrockLLM, extract_provider
|
||||
from mem0.utils.factory import LlmFactory
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def mock_boto3():
|
||||
"""Patch boto3 so no real AWS calls are made during unit tests."""
|
||||
with patch("mem0.llms.aws_bedrock.boto3") as mock_b3:
|
||||
runtime_client = MagicMock()
|
||||
bedrock_client = MagicMock()
|
||||
bedrock_client.list_foundation_models.return_value = {"modelSummaries": []}
|
||||
|
||||
def _client(service, **kwargs):
|
||||
if service == "bedrock-runtime":
|
||||
return runtime_client
|
||||
return bedrock_client
|
||||
|
||||
mock_b3.client.side_effect = _client
|
||||
yield runtime_client
|
||||
|
||||
|
||||
def _make_llm(model: str, mock_boto3, **kwargs) -> AWSBedrockLLM:
|
||||
"""Instantiate AWSBedrockLLM with a given model, all AWS calls mocked."""
|
||||
config = AWSBedrockConfig(model=model, **kwargs)
|
||||
return AWSBedrockLLM(config)
|
||||
|
||||
|
||||
def _converse_response(text: str = "ok") -> dict:
|
||||
"""Minimal Converse API response dict."""
|
||||
return {"output": {"message": {"content": [{"text": text}]}}}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# extract_provider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestExtractProvider:
|
||||
def test_standard_anthropic_model(self):
|
||||
assert extract_provider("anthropic.claude-3-5-sonnet-20240620-v1:0") == "anthropic"
|
||||
|
||||
def test_inference_profile_us_prefix(self):
|
||||
# Cross-region inference profile IDs look like us.anthropic.<model>
|
||||
assert extract_provider("us.anthropic.claude-haiku-4-5-20251001-v1:0") == "anthropic"
|
||||
|
||||
def test_inference_profile_eu_prefix(self):
|
||||
assert extract_provider("eu.anthropic.claude-sonnet-4-5-20250929-v1:0") == "anthropic"
|
||||
|
||||
def test_inference_profile_ap_prefix(self):
|
||||
assert extract_provider("ap.anthropic.claude-3-opus-20240229-v1:0") == "anthropic"
|
||||
|
||||
def test_amazon_model(self):
|
||||
assert extract_provider("amazon.nova-3-mini-20241119-v1:0") == "amazon"
|
||||
|
||||
def test_meta_model(self):
|
||||
assert extract_provider("meta.llama3-8b-instruct-v1:0") == "meta"
|
||||
|
||||
def test_mistral_model(self):
|
||||
assert extract_provider("mistral.mistral-7b-instruct-v0:2") == "mistral"
|
||||
|
||||
def test_unknown_model_raises(self):
|
||||
with pytest.raises(ValueError, match="Unknown provider"):
|
||||
extract_provider("unknown-vendor.some-model-v1:0")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AWSBedrockConfig
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAWSBedrockConfig:
|
||||
def test_top_p_defaults_to_none(self):
|
||||
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0")
|
||||
assert config.top_p is None
|
||||
|
||||
def test_top_p_explicit_value_stored(self):
|
||||
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", top_p=0.8)
|
||||
assert config.top_p == 0.8
|
||||
|
||||
def test_get_model_config_excludes_top_p_by_default(self):
|
||||
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", temperature=0.5)
|
||||
model_cfg = config.get_model_config()
|
||||
assert "top_p" not in model_cfg
|
||||
|
||||
def test_get_model_config_includes_top_p_when_set(self):
|
||||
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", top_p=0.7)
|
||||
model_cfg = config.get_model_config()
|
||||
assert model_cfg["top_p"] == 0.7
|
||||
|
||||
def test_get_model_config_top_p_via_model_kwargs(self):
|
||||
"""model_kwargs can supply top_p after merge; same semantics as explicit top_p."""
|
||||
config = AWSBedrockConfig(
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
model_kwargs={"top_p": 0.88},
|
||||
)
|
||||
assert config.get_model_config()["top_p"] == 0.88
|
||||
|
||||
def test_get_model_config_always_includes_temperature(self):
|
||||
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", temperature=0.3)
|
||||
model_cfg = config.get_model_config()
|
||||
assert model_cfg["temperature"] == 0.3
|
||||
|
||||
def test_aws_region_stored(self):
|
||||
config = AWSBedrockConfig(
|
||||
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
aws_region="us-east-2",
|
||||
)
|
||||
assert config.aws_region == "us-east-2"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LlmFactory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestLlmFactory:
|
||||
def test_aws_bedrock_uses_aws_bedrock_config(self):
|
||||
_, config_class = LlmFactory.provider_to_class["aws_bedrock"]
|
||||
assert config_class is AWSBedrockConfig
|
||||
|
||||
def test_factory_create_accepts_aws_region(self, mock_boto3):
|
||||
"""LlmFactory.create must not crash when aws_region is in the config dict.
|
||||
|
||||
Before the fix, the factory mapped aws_bedrock to BaseLlmConfig which has
|
||||
no aws_region parameter, causing: TypeError: __init__() got an unexpected
|
||||
keyword argument 'aws_region'.
|
||||
"""
|
||||
user_dict = {
|
||||
"model": "anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
"aws_region": "us-east-2",
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 2000,
|
||||
}
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
|
||||
# This must not raise TypeError
|
||||
llm = LlmFactory.create("aws_bedrock", user_dict)
|
||||
|
||||
assert isinstance(llm.config, AWSBedrockConfig)
|
||||
assert llm.config.aws_region == "us-east-2"
|
||||
assert llm.config.top_p is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _build_inference_config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestBuildInferenceConfig:
|
||||
"""
|
||||
Unit tests for the _build_inference_config helper.
|
||||
Validates the exact keys present in the returned dict.
|
||||
"""
|
||||
|
||||
def test_anthropic_only_temperature_by_default(self, mock_boto3):
|
||||
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.5)
|
||||
cfg = llm._build_inference_config()
|
||||
assert "temperature" in cfg
|
||||
assert cfg["temperature"] == 0.5
|
||||
assert "topP" not in cfg, "topP must be absent when top_p not configured"
|
||||
|
||||
def test_anthropic_top_p_explicitly_set_still_omits_top_p(self, mock_boto3):
|
||||
# Anthropic rejects both; topP must still be omitted even when set
|
||||
llm = _make_llm(
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
mock_boto3,
|
||||
temperature=0.5,
|
||||
top_p=0.9,
|
||||
)
|
||||
cfg = llm._build_inference_config()
|
||||
assert "temperature" in cfg
|
||||
assert "topP" not in cfg
|
||||
|
||||
def test_anthropic_inference_profile_omits_top_p(self, mock_boto3):
|
||||
# Cross-region inference profiles (us.anthropic.*) follow the same rule
|
||||
llm = _make_llm("us.anthropic.claude-haiku-4-5-20251001-v1:0", mock_boto3, temperature=0.1)
|
||||
cfg = llm._build_inference_config()
|
||||
assert "topP" not in cfg
|
||||
|
||||
def test_amazon_includes_top_p_when_set(self, mock_boto3):
|
||||
llm = _make_llm(
|
||||
"amazon.nova-3-mini-20241119-v1:0",
|
||||
mock_boto3,
|
||||
temperature=0.5,
|
||||
top_p=0.85,
|
||||
)
|
||||
cfg = llm._build_inference_config()
|
||||
assert cfg["temperature"] == 0.5
|
||||
assert cfg["topP"] == 0.85
|
||||
|
||||
def test_amazon_omits_top_p_when_not_set(self, mock_boto3):
|
||||
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3, temperature=0.5)
|
||||
cfg = llm._build_inference_config()
|
||||
assert "topP" not in cfg
|
||||
|
||||
def test_max_tokens_present(self, mock_boto3):
|
||||
llm = _make_llm(
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
mock_boto3,
|
||||
max_tokens=1024,
|
||||
)
|
||||
cfg = llm._build_inference_config()
|
||||
assert cfg["maxTokens"] == 1024
|
||||
|
||||
def test_nova_fallback_max_tokens_when_absent(self, mock_boto3):
|
||||
"""Legacy Nova Converse used 5000 when max_tokens was missing from the dict."""
|
||||
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3)
|
||||
llm.model_config.pop("max_tokens", None)
|
||||
cfg = llm._build_inference_config()
|
||||
assert cfg["maxTokens"] == 5000
|
||||
|
||||
def test_anthropic_fallback_max_tokens_when_absent(self, mock_boto3):
|
||||
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3)
|
||||
llm.model_config.pop("max_tokens", None)
|
||||
cfg = llm._build_inference_config()
|
||||
assert cfg["maxTokens"] == 2000
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# generate_response — Converse API call assertions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
MESSAGES = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
||||
TOOLS = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "add_memory",
|
||||
"description": "Store a memory",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"data": {"type": "string"}},
|
||||
"required": ["data"],
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
class TestGenerateResponseConverse:
|
||||
def test_standard_anthropic_no_top_p_in_converse_call(self, mock_boto3):
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.2)
|
||||
|
||||
llm.generate_response(MESSAGES)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
inference_cfg = kwargs["inferenceConfig"]
|
||||
assert "topP" not in inference_cfg
|
||||
assert inference_cfg["temperature"] == 0.2
|
||||
|
||||
def test_standard_anthropic_inference_profile_no_top_p(self, mock_boto3):
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm("us.anthropic.claude-haiku-4-5-20251001-v1:0", mock_boto3, temperature=0.1)
|
||||
|
||||
llm.generate_response(MESSAGES)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert "topP" not in kwargs["inferenceConfig"]
|
||||
|
||||
def test_with_tools_anthropic_no_top_p(self, mock_boto3):
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.2)
|
||||
|
||||
llm.generate_response(MESSAGES, tools=TOOLS)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert "topP" not in kwargs["inferenceConfig"]
|
||||
|
||||
def test_with_tools_anthropic_top_p_set_still_omitted(self, mock_boto3):
|
||||
# Even if user explicitly sets top_p, Anthropic must not receive topP
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm(
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
mock_boto3,
|
||||
temperature=0.2,
|
||||
top_p=0.9,
|
||||
)
|
||||
|
||||
llm.generate_response(MESSAGES, tools=TOOLS)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert "topP" not in kwargs["inferenceConfig"]
|
||||
|
||||
def test_nova_includes_top_p_when_explicitly_set(self, mock_boto3):
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm(
|
||||
"amazon.nova-3-mini-20241119-v1:0",
|
||||
mock_boto3,
|
||||
temperature=0.5,
|
||||
top_p=0.85,
|
||||
)
|
||||
|
||||
llm.generate_response(MESSAGES)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert kwargs["inferenceConfig"]["topP"] == 0.85
|
||||
|
||||
def test_nova_omits_top_p_when_not_set(self, mock_boto3):
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3, temperature=0.5)
|
||||
|
||||
llm.generate_response(MESSAGES)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert "topP" not in kwargs["inferenceConfig"]
|
||||
|
||||
def test_anthropic_model_kwargs_top_p_still_omits_top_p_in_converse(self, mock_boto3):
|
||||
"""top_p injected via model_kwargs must not add topP for Anthropic Converse."""
|
||||
mock_boto3.converse.return_value = _converse_response()
|
||||
llm = _make_llm(
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
mock_boto3,
|
||||
temperature=0.2,
|
||||
model_kwargs={"top_p": 0.88},
|
||||
)
|
||||
|
||||
llm.generate_response(MESSAGES)
|
||||
|
||||
_, kwargs = mock_boto3.converse.call_args
|
||||
assert "topP" not in kwargs["inferenceConfig"]
|
||||
Reference in New Issue
Block a user