From 9ead7175313ab55f6323dc6b06c378d4335fdbe5 Mon Sep 17 00:00:00 2001 From: aswny <87371411+aswny@users.noreply.github.com> Date: Thu, 25 Apr 2024 17:19:55 +0000 Subject: [PATCH 1/2] fix Llama models message to prompt conversion in for AWS Bedrock provider --- litellm/llms/bedrock.py | 4 ++++ litellm/llms/prompt_templates/factory.py | 10 ++++++++++ 2 files changed, 14 insertions(+) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index ef6dbfb1b8..149b684724 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -653,6 +653,10 @@ def convert_messages_to_prompt(model, messages, provider, custom_prompt_dict): prompt = prompt_factory( model=model, messages=messages, custom_llm_provider="bedrock" ) + elif provider == "meta": + prompt = prompt_factory( + model=model, messages=messages, custom_llm_provider="bedrock" + ) else: prompt = "" for message in messages: diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index a6d1d64386..00c6229c9d 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -1346,6 +1346,16 @@ def prompt_factory( return anthropic_pt(messages=messages) elif "mistral." in model: return mistral_instruct_pt(messages=messages) + elif "llama2" in model: + return llama_2_chat_pt(messages=messages) + elif "llama3" in model: + return hf_chat_template( + model=model, + messages=messages, + chat_template=known_tokenizer_config[ # type: ignore + "meta-llama/Meta-Llama-3-8B-Instruct" + ]["tokenizer"]["chat_template"], + ) elif custom_llm_provider == "perplexity": for message in messages: message.pop("name", None) From 781af56f485c0d6b3b33aa71cb34df125e3e7103 Mon Sep 17 00:00:00 2001 From: aswny <87371411+aswny@users.noreply.github.com> Date: Thu, 25 Apr 2024 17:52:38 +0000 Subject: [PATCH 2/2] check model type chat/instruct to apply template --- 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 00c6229c9d..ed719f600d 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -1346,9 +1346,9 @@ def prompt_factory( return anthropic_pt(messages=messages) elif "mistral." in model: return mistral_instruct_pt(messages=messages) - elif "llama2" in model: + elif "llama2" in model and "chat" in model: return llama_2_chat_pt(messages=messages) - elif "llama3" in model: + elif "llama3" in model and "instruct" in model: return hf_chat_template( model=model, messages=messages,