From 739f4f05f60a75a99d2ec7aa3989fea975f5db57 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 22:45:54 -0500 Subject: [PATCH 01/12] add support for bedrock mistral models --- litellm/llms/bedrock.py | 14 +++++ litellm/llms/prompt_templates/factory.py | 2 + litellm/tests/test_bedrock_completion.py | 80 +++++++++++++++++------- 3 files changed, 72 insertions(+), 24 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index b7f1c502368..4806a57e2a8 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -492,6 +492,8 @@ def convert_messages_to_prompt(model, messages, provider, custom_prompt_dict): prompt = prompt_factory( model=model, messages=messages, custom_llm_provider="bedrock" ) + elif provider == "mistral": + prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") else: prompt = "" for message in messages: @@ -623,7 +625,16 @@ def completion( "textGenerationConfig": inference_params, } ) + elif provider == "mistral": + ## LOAD CONFIG + config = litellm.AmazonLlamaConfig.get_config() + for k, v in config.items(): + if ( + k not in inference_params + ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in + inference_params[k] = v + data = json.dumps({"prompt": prompt, **inference_params}) else: data = json.dumps({}) @@ -729,6 +740,9 @@ def completion( outputText = response_body["generations"][0]["text"] elif provider == "meta": outputText = response_body["generation"] + elif provider == "mistral": + outputText = response_body["outputs"][0]["text"] + model_response["finish_reason"] = response_body["outputs"][0]["stop_reason"] else: # amazon titan outputText = response_body.get("results")[0].get("outputText") diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index 4ed4d9295de..cc8d3d49bd8 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -674,6 +674,8 @@ def prompt_factory( return claude_2_1_pt(messages=messages) else: return anthropic_pt(messages=messages) + elif "mistral." in model: + return mistral_instruct_pt(messages=messages) try: if "meta-llama/llama-2" in model and "chat" in model: return llama_2_chat_pt(messages=messages) diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index a448fc3a57f..4a9164019e2 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -1,33 +1,33 @@ # @pytest.mark.skip(reason="AWS Suspended Account") -# import sys, os -# import traceback -# from dotenv import load_dotenv +import sys, os +import traceback +from dotenv import load_dotenv -# load_dotenv() -# import os, io +load_dotenv() +import os, io -# sys.path.insert( -# 0, os.path.abspath("../..") -# ) # Adds the parent directory to the system path -# import pytest -# import litellm -# from litellm import embedding, completion, completion_cost, Timeout -# from litellm import RateLimitError +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path +import pytest +import litellm +from litellm import embedding, completion, completion_cost, Timeout +from litellm import RateLimitError -# # litellm.num_retries = 3 -# litellm.cache = None -# litellm.success_callback = [] -# user_message = "Write a short poem about the sky" -# messages = [{"content": user_message, "role": "user"}] +# litellm.num_retries = 3 +litellm.cache = None +litellm.success_callback = [] +user_message = "Write a short poem about the sky" +messages = [{"content": user_message, "role": "user"}] -# @pytest.fixture(autouse=True) -# def reset_callbacks(): -# print("\npytest fixture - resetting callbacks") -# litellm.success_callback = [] -# litellm._async_success_callback = [] -# litellm.failure_callback = [] -# litellm.callbacks = [] +@pytest.fixture(autouse=True) +def reset_callbacks(): + print("\npytest fixture - resetting callbacks") + litellm.success_callback = [] + litellm._async_success_callback = [] + litellm.failure_callback = [] + litellm.callbacks = [] # def test_completion_bedrock_claude_completion_auth(): @@ -257,3 +257,35 @@ # # test_provisioned_throughput() + +def test_completion_bedrock_mistral_completion_auth(): + print("calling bedrock mistral completion params auth") + import os + + # aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"] + # aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"] + # aws_region_name = os.environ["AWS_REGION_NAME"] + + # os.environ.pop("AWS_ACCESS_KEY_ID", None) + # os.environ.pop("AWS_SECRET_ACCESS_KEY", None) + # os.environ.pop("AWS_REGION_NAME", None) + try: + response = completion( + model="bedrock/mistral.mistral-7b-instruct-v0:2", + messages=messages, + max_tokens=10, + temperature=0.1, + ) + # Add any assertions here to check the response + print(response) + + # os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id + # os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key + # os.environ["AWS_REGION_NAME"] = aws_region_name + except RateLimitError: + pass + except Exception as e: + pytest.fail(f"Error occurred: {e}") + + +test_completion_bedrock_mistral_completion_auth() \ No newline at end of file From 8907b2733a2567cc4881fe91740c596e3f01fec1 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 22:49:31 -0500 Subject: [PATCH 02/12] skip test but it did work locally --- litellm/tests/test_bedrock_completion.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index 4a9164019e2..e9ad6ac1c34 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -257,7 +257,7 @@ def reset_callbacks(): # # test_provisioned_throughput() - +@pytest.mark.skip(reason="AWS Suspended Account") def test_completion_bedrock_mistral_completion_auth(): print("calling bedrock mistral completion params auth") import os From b4dc7f0f17ea2d3f0cc6f8f1d1153f0cb4347ee1 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:14:00 -0500 Subject: [PATCH 03/12] Add AmazonMistralConfig --- litellm/__init__.py | 1 + litellm/llms/bedrock.py | 51 ++++++++++++++++++++++++++++++++++++++++- 2 files changed, 51 insertions(+), 1 deletion(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index cd639ddb9b7..65807a8a084 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -591,6 +591,7 @@ from .llms.bedrock import ( AmazonCohereConfig, AmazonLlamaConfig, AmazonStabilityConfig, + AmazonMistralConfig ) from .llms.openai import OpenAIConfig, OpenAITextCompletionConfig from .llms.azure import AzureOpenAIConfig, AzureOpenAIError diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 4806a57e2a8..3f14ac9e410 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -282,6 +282,55 @@ class AmazonLlamaConfig: } +class AmazonMistralConfig: + """ + Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-mistral.html + Supported Params for the Amazon / Mistral models: + + - `max_tokens` (integer) max tokens, + - `temperature` (float) temperature for model, + - `top_p` (float) top p for model + - `stop` [string] A list of stop sequences that if generated by the model, stops the model from generating further output. + - `top_k` (float) top k for model + """ + + max_tokens: Optional[int] = None + temperature: Optional[float] = None + topP: Optional[float] = None + topK: Optional[float] = None + stop: Optional[list[str]] = None + + def __init__( + self, + max_tokens: Optional[int] = None, + temperature: Optional[float] = None, + topP: Optional[int] = None, + topK: Optional[float] = None, + stop: Optional[list[str]] = None, + ) -> None: + locals_ = locals() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return { + k: v + for k, v in cls.__dict__.items() + if not k.startswith("__") + and not isinstance( + v, + ( + types.FunctionType, + types.BuiltinFunctionType, + classmethod, + staticmethod, + ), + ) + and v is not None + } + class AmazonStabilityConfig: """ Reference: https://us-west-2.console.aws.amazon.com/bedrock/home?region=us-west-2#/providers?model=stability.stable-diffusion-xl-v0 @@ -627,7 +676,7 @@ def completion( ) elif provider == "mistral": ## LOAD CONFIG - config = litellm.AmazonLlamaConfig.get_config() + config = litellm.AmazonMistralConfig.get_config() for k, v in config.items(): if ( k not in inference_params From 58ed6e77de83f3f3e7f74947c8e4195173553870 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:27:02 -0500 Subject: [PATCH 04/12] add assertion for test --- litellm/tests/test_bedrock_completion.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index e9ad6ac1c34..6f322207428 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -11,7 +11,7 @@ sys.path.insert( ) # Adds the parent directory to the system path import pytest import litellm -from litellm import embedding, completion, completion_cost, Timeout +from litellm import embedding, completion, completion_cost, Timeout, ModelResponse from litellm import RateLimitError # litellm.num_retries = 3 @@ -270,14 +270,15 @@ def test_completion_bedrock_mistral_completion_auth(): # os.environ.pop("AWS_SECRET_ACCESS_KEY", None) # os.environ.pop("AWS_REGION_NAME", None) try: - response = completion( + response:ModelResponse = completion( model="bedrock/mistral.mistral-7b-instruct-v0:2", messages=messages, max_tokens=10, temperature=0.1, ) # Add any assertions here to check the response - print(response) + assert len(response.choices) > 0 + assert len(response.choices[0].message.content) > 0 # os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id # os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key From 2321f19fe7268e82f3eecf9e82aad1fbefcde79b Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:28:25 -0500 Subject: [PATCH 05/12] comment out tests --- litellm/tests/test_bedrock_completion.py | 121 ++++++++++++----------- 1 file changed, 61 insertions(+), 60 deletions(-) diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index 6f322207428..3e3d8b6bbcd 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -1,33 +1,33 @@ # @pytest.mark.skip(reason="AWS Suspended Account") -import sys, os -import traceback -from dotenv import load_dotenv - -load_dotenv() -import os, io - -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path -import pytest -import litellm -from litellm import embedding, completion, completion_cost, Timeout, ModelResponse -from litellm import RateLimitError - -# litellm.num_retries = 3 -litellm.cache = None -litellm.success_callback = [] -user_message = "Write a short poem about the sky" -messages = [{"content": user_message, "role": "user"}] - - -@pytest.fixture(autouse=True) -def reset_callbacks(): - print("\npytest fixture - resetting callbacks") - litellm.success_callback = [] - litellm._async_success_callback = [] - litellm.failure_callback = [] - litellm.callbacks = [] +# import sys, os +# import traceback +# from dotenv import load_dotenv +# +# load_dotenv() +# import os, io +# +# sys.path.insert( +# 0, os.path.abspath("../..") +# ) # Adds the parent directory to the system path +# import pytest +# import litellm +# from litellm import embedding, completion, completion_cost, Timeout, ModelResponse +# from litellm import RateLimitError +# +# # litellm.num_retries = 3 +# litellm.cache = None +# litellm.success_callback = [] +# user_message = "Write a short poem about the sky" +# messages = [{"content": user_message, "role": "user"}] +# +# +# @pytest.fixture(autouse=True) +# def reset_callbacks(): +# print("\npytest fixture - resetting callbacks") +# litellm.success_callback = [] +# litellm._async_success_callback = [] +# litellm.failure_callback = [] +# litellm.callbacks = [] # def test_completion_bedrock_claude_completion_auth(): @@ -257,36 +257,37 @@ def reset_callbacks(): # # test_provisioned_throughput() -@pytest.mark.skip(reason="AWS Suspended Account") -def test_completion_bedrock_mistral_completion_auth(): - print("calling bedrock mistral completion params auth") - import os - - # aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"] - # aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"] - # aws_region_name = os.environ["AWS_REGION_NAME"] - - # os.environ.pop("AWS_ACCESS_KEY_ID", None) - # os.environ.pop("AWS_SECRET_ACCESS_KEY", None) - # os.environ.pop("AWS_REGION_NAME", None) - try: - response:ModelResponse = completion( - model="bedrock/mistral.mistral-7b-instruct-v0:2", - messages=messages, - max_tokens=10, - temperature=0.1, - ) - # Add any assertions here to check the response - assert len(response.choices) > 0 - assert len(response.choices[0].message.content) > 0 - - # os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id - # os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key - # os.environ["AWS_REGION_NAME"] = aws_region_name - except RateLimitError: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") -test_completion_bedrock_mistral_completion_auth() \ No newline at end of file +# def test_completion_bedrock_mistral_completion_auth(): +# print("calling bedrock mistral completion params auth") +# import os +# +# # aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"] +# # aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"] +# # aws_region_name = os.environ["AWS_REGION_NAME"] +# +# # os.environ.pop("AWS_ACCESS_KEY_ID", None) +# # os.environ.pop("AWS_SECRET_ACCESS_KEY", None) +# # os.environ.pop("AWS_REGION_NAME", None) +# try: +# response:ModelResponse = completion( +# model="bedrock/mistral.mistral-7b-instruct-v0:2", +# messages=messages, +# max_tokens=10, +# temperature=0.1, +# ) +# # Add any assertions here to check the response +# assert len(response.choices) > 0 +# assert len(response.choices[0].message.content) > 0 +# +# # os.environ["AWS_ACCESS_KEY_ID"] = aws_access_key_id +# # os.environ["AWS_SECRET_ACCESS_KEY"] = aws_secret_access_key +# # os.environ["AWS_REGION_NAME"] = aws_region_name +# except RateLimitError: +# pass +# except Exception as e: +# pytest.fail(f"Error occurred: {e}") +# +# +# test_completion_bedrock_mistral_completion_auth() \ No newline at end of file From dc84e28c41b23f7fd30f40a76f8a3b0fd10dbc43 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:34:04 -0500 Subject: [PATCH 06/12] update docs for supported models --- docs/my-website/docs/providers/bedrock.md | 20 +++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index b5705718445..d744596c4fc 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -286,18 +286,20 @@ response = litellm.embedding( ## Supported AWS Bedrock Models Here's an example of using a bedrock model with LiteLLM -| Model Name | Command | -|--------------------------|------------------------------------------------------------------| +| Model Name | Command | +|----------------------------|------------------------------------------------------------------| | Anthropic Claude-V2.1 | `completion(model='bedrock/anthropic.claude-v2:1', messages=messages)` | `os.environ['ANTHROPIC_ACCESS_KEY_ID']`, `os.environ['ANTHROPIC_SECRET_ACCESS_KEY']` | -| Anthropic Claude-V2 | `completion(model='bedrock/anthropic.claude-v2', messages=messages)` | `os.environ['ANTHROPIC_ACCESS_KEY_ID']`, `os.environ['ANTHROPIC_SECRET_ACCESS_KEY']` | +| Anthropic Claude-V2 | `completion(model='bedrock/anthropic.claude-v2', messages=messages)` | `os.environ['ANTHROPIC_ACCESS_KEY_ID']`, `os.environ['ANTHROPIC_SECRET_ACCESS_KEY']` | | Anthropic Claude-Instant V1 | `completion(model='bedrock/anthropic.claude-instant-v1', messages=messages)` | `os.environ['ANTHROPIC_ACCESS_KEY_ID']`, `os.environ['ANTHROPIC_SECRET_ACCESS_KEY']` | -| Amazon Titan Lite | `completion(model='bedrock/amazon.titan-text-lite-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | -| Amazon Titan Express | `completion(model='bedrock/amazon.titan-text-express-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | -| Cohere Command | `completion(model='bedrock/cohere.command-text-v14', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | -| AI21 J2-Mid | `completion(model='bedrock/ai21.j2-mid-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Amazon Titan Lite | `completion(model='bedrock/amazon.titan-text-lite-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Amazon Titan Express | `completion(model='bedrock/amazon.titan-text-express-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Cohere Command | `completion(model='bedrock/cohere.command-text-v14', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| AI21 J2-Mid | `completion(model='bedrock/ai21.j2-mid-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | | AI21 J2-Ultra | `completion(model='bedrock/ai21.j2-ultra-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | -| Meta Llama 2 Chat 13b | `completion(model='bedrock/meta.llama2-13b-chat-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | -| Meta Llama 2 Chat 70b | `completion(model='bedrock/meta.llama2-70b-chat-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Meta Llama 2 Chat 13b | `completion(model='bedrock/meta.llama2-13b-chat-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Meta Llama 2 Chat 70b | `completion(model='bedrock/meta.llama2-70b-chat-v1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Mistral 7B Instruct | `completion(model='bedrock/mistral.mistral-7b-instruct-v0:2', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Mixtral 8x7B Instruct | `completion(model='bedrock/mistral.mixtral-8x7b-instruct-v0:1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | ## Bedrock Embedding From bf92f84b001391f8c7f97da924231f46a38b7591 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:53:16 -0500 Subject: [PATCH 07/12] update prices and context window --- model_prices_and_context_window.json | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 66061acc4e9..13e54bdf99d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1211,6 +1211,20 @@ "litellm_provider": "bedrock", "mode": "embedding" }, + "bedrock/us-west-2/mistral.mixtral-8x7b-instruct": { + "max_tokens": 32000, + "input_cost_per_token": 0.00000045, + "output_cost_per_token": 0.0000007, + "litellm_provider": "bedrock", + "mode": "completion" + }, + "bedrock/us-west-2/mistral.mistral-7b-instruct": { + "max_tokens": 32000, + "input_cost_per_token": 0.00000015, + "output_cost_per_token": 0.0000002, + "litellm_provider": "bedrock", + "mode": "completion" + }, "anthropic.claude-v1": { "max_tokens": 100000, "max_output_tokens": 8191, @@ -2195,4 +2209,4 @@ "mode": "embedding" } -} +} \ No newline at end of file From 2c97b952eae2bc90ec921dcb284103c0833ba528 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Fri, 1 Mar 2024 23:58:03 -0500 Subject: [PATCH 08/12] follow camelcase convention --- litellm/llms/bedrock.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 3f14ac9e410..e0938ded8ee 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -287,14 +287,14 @@ class AmazonMistralConfig: Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-mistral.html Supported Params for the Amazon / Mistral models: - - `max_tokens` (integer) max tokens, + - `maxTokens` (integer) max tokens, - `temperature` (float) temperature for model, - - `top_p` (float) top p for model + - `topP` (float) top p for model - `stop` [string] A list of stop sequences that if generated by the model, stops the model from generating further output. - - `top_k` (float) top k for model + - `topK` (float) top k for model """ - max_tokens: Optional[int] = None + maxTokens: Optional[int] = None temperature: Optional[float] = None topP: Optional[float] = None topK: Optional[float] = None @@ -302,7 +302,7 @@ class AmazonMistralConfig: def __init__( self, - max_tokens: Optional[int] = None, + maxTokens: Optional[int] = None, temperature: Optional[float] = None, topP: Optional[int] = None, topK: Optional[float] = None, @@ -1118,4 +1118,4 @@ def image_generation( image_dict = {"url": artifact["base64"]} model_response.data = image_dict - return model_response + return model_response \ No newline at end of file From 67110df3e13ed1b1a0e1d5d825aa1ad5262dd448 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Sat, 2 Mar 2024 00:07:21 -0500 Subject: [PATCH 09/12] change to snake case 'cause of aws docs --- litellm/llms/bedrock.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index e0938ded8ee..1a25e167a9d 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -287,25 +287,25 @@ class AmazonMistralConfig: Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-mistral.html Supported Params for the Amazon / Mistral models: - - `maxTokens` (integer) max tokens, + - `max_tokens` (integer) max tokens, - `temperature` (float) temperature for model, - - `topP` (float) top p for model + - `top_p` (float) top p for model - `stop` [string] A list of stop sequences that if generated by the model, stops the model from generating further output. - - `topK` (float) top k for model + - `top_k` (float) top k for model """ - maxTokens: Optional[int] = None + max_tokens: Optional[int] = None temperature: Optional[float] = None - topP: Optional[float] = None - topK: Optional[float] = None + top_p: Optional[float] = None + top_k: Optional[float] = None stop: Optional[list[str]] = None def __init__( self, - maxTokens: Optional[int] = None, + max_tokens: Optional[int] = None, temperature: Optional[float] = None, - topP: Optional[int] = None, - topK: Optional[float] = None, + top_p: Optional[int] = None, + top_k: Optional[float] = None, stop: Optional[list[str]] = None, ) -> None: locals_ = locals() From a4e24761a037624b11d4d565c3ad588e428d7286 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Sat, 2 Mar 2024 13:25:04 -0500 Subject: [PATCH 10/12] map optional params --- litellm/utils.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/litellm/utils.py b/litellm/utils.py index 6d64368f879..b0cca2cb397 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4551,6 +4551,21 @@ def get_optional_params( optional_params["temperature"] = temperature if max_tokens is not None: optional_params["max_tokens"] = max_tokens + elif "mistral" in model: + supported_params = ["max_tokens", "temperature", "stop", "top_p", "stream"] + _check_valid_arg(supported_params=supported_params) + # mistral params on bedrock + # \"max_tokens_to_sample\":300,\"temperature\":0.5,\"top_p\":1,\"stop_sequences\":[\"\\\\n\\\\nHuman:\"]}" + if max_tokens is not None: + optional_params["max_tokens"] = max_tokens + if temperature is not None: + optional_params["temperature"] = temperature + if top_p is not None: + optional_params["top_p"] = top_p + if stop is not None: + optional_params["stop"] = stop + if stream is not None: + optional_params["stream"] = stream elif custom_llm_provider == "aleph_alpha": supported_params = [ "max_tokens", @@ -9657,4 +9672,4 @@ def _get_base_model_from_metadata(model_call_details=None): base_model = model_info.get("base_model", None) if base_model is not None: return base_model - return None + return None \ No newline at end of file From 12d7ea914a9e988ee2d40c3e284e0c26865dc7d3 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Sat, 2 Mar 2024 13:34:39 -0500 Subject: [PATCH 11/12] update comments --- litellm/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/utils.py b/litellm/utils.py index b0cca2cb397..67bf7108f2c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4555,7 +4555,7 @@ def get_optional_params( supported_params = ["max_tokens", "temperature", "stop", "top_p", "stream"] _check_valid_arg(supported_params=supported_params) # mistral params on bedrock - # \"max_tokens_to_sample\":300,\"temperature\":0.5,\"top_p\":1,\"stop_sequences\":[\"\\\\n\\\\nHuman:\"]}" + # \"max_tokens\":400,\"temperature\":0.7,\"top_p\":0.7,\"stop\":[\"\\\\n\\\\nHuman:\"]}" if max_tokens is not None: optional_params["max_tokens"] = max_tokens if temperature is not None: From 446cc393f92fffe77833eba47bc32f357188a582 Mon Sep 17 00:00:00 2001 From: Tim Xia Date: Sat, 2 Mar 2024 22:27:42 -0500 Subject: [PATCH 12/12] Fix message template based on transformers chat_template docs --- litellm/llms/prompt_templates/factory.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index cc8d3d49bd8..103eb5977ac 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -110,9 +110,9 @@ def mistral_instruct_pt(messages): "post_message": " [/INST]\n", }, "user": {"pre_message": "[INST] ", "post_message": " [/INST]\n"}, - "assistant": {"pre_message": " ", "post_message": " "}, + "assistant": {"pre_message": " ", "post_message": " "}, }, - final_prompt_value="", + final_prompt_value="", messages=messages, ) return prompt