From 0db7fa3fd8d5ec4140b5e34154b5819280737c9c Mon Sep 17 00:00:00 2001 From: alisalim17 Date: Mon, 29 Apr 2024 14:20:24 +0400 Subject: [PATCH 1/2] fix: cohere tool results --- litellm/llms/prompt_templates/factory.py | 80 ++++++++++++++++++------ 1 file changed, 61 insertions(+), 19 deletions(-) diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index c51dc89be5..e1fa354c63 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -3,8 +3,14 @@ import requests, traceback import json, re, xml.etree.ElementTree as ET from jinja2 import Template, exceptions, meta, BaseLoader from jinja2.sandbox import ImmutableSandboxedEnvironment -from typing import Optional, Any -from typing import List +from typing import ( + Any, + List, + Mapping, + MutableMapping, + Optional, + Sequence, +) import litellm @@ -430,8 +436,10 @@ def format_prompt_togetherai(messages, prompt_format, chat_template): prompt = default_pt(messages) return prompt + ### IBM Granite + def ibm_granite_pt(messages: list): """ IBM's Granite models uses the template: @@ -440,23 +448,24 @@ def ibm_granite_pt(messages: list): See: https://www.ibm.com/docs/en/watsonx-as-a-service?topic=solutions-supported-foundation-models """ return custom_prompt( - messages=messages, + messages=messages, role_dict={ - 'system': { - 'pre_message': '<|system|>\n', - 'post_message': '\n', + "system": { + "pre_message": "<|system|>\n", + "post_message": "\n", }, - 'user': { - 'pre_message': '<|user|>\n', - 'post_message': '\n', + "user": { + "pre_message": "<|user|>\n", + "post_message": "\n", }, - 'assistant': { - 'pre_message': '<|assistant|>\n', - 'post_message': '\n', - } - } + "assistant": { + "pre_message": "<|assistant|>\n", + "post_message": "\n", + }, + }, ).strip() + ### ANTHROPIC ### @@ -1043,6 +1052,30 @@ def get_system_prompt(messages): return system_prompt, messages +def convert_to_documents( + observations: Any, +) -> List[MutableMapping]: + """Converts observations into a 'document' dict""" + documents: List[MutableMapping] = [] + if isinstance(observations, str): + # strings are turned into a key/value pair and a key of 'output' is added. + observations = [{"output": observations}] + elif isinstance(observations, Mapping): + # single mappings are transformed into a list to simplify the rest of the code. + observations = [observations] + elif not isinstance(observations, Sequence): + # all other types are turned into a key/value pair within a list + observations = [{"output": observations}] + + for doc in observations: + if not isinstance(doc, Mapping): + # types that aren't Mapping are turned into a key/value pair. + doc = {"output": doc} + documents.append(doc) + + return documents + + def convert_openai_message_to_cohere_tool_result(message): """ OpenAI message with a tool result looks like: @@ -1084,7 +1117,7 @@ def convert_openai_message_to_cohere_tool_result(message): "parameters": {"location": "San Francisco, CA"}, "generation_id": tool_call_id, }, - "outputs": [content], + "outputs": convert_to_documents(content), } return cohere_tool_result @@ -1097,7 +1130,7 @@ def cohere_message_pt(messages: list): if message["role"] == "tool": tool_result = convert_openai_message_to_cohere_tool_result(message) tool_results.append(tool_result) - else: + elif message.get("content"): prompt += message["content"] + "\n\n" prompt = prompt.rstrip() return prompt, tool_results @@ -1396,9 +1429,18 @@ def prompt_factory( # https://llama.meta.com/docs/model-cards-and-prompt-formats/meta-llama-3/ return custom_prompt( role_dict={ - "system": {"pre_message": "<|start_header_id|>system<|end_header_id|>\n", "post_message": "<|eot_id|>"}, - "user": {"pre_message": "<|start_header_id|>user<|end_header_id|>\n", "post_message": "<|eot_id|>"}, - "assistant": {"pre_message": "<|start_header_id|>assistant<|end_header_id|>\n", "post_message": "<|eot_id|>"}, + "system": { + "pre_message": "<|start_header_id|>system<|end_header_id|>\n", + "post_message": "<|eot_id|>", + }, + "user": { + "pre_message": "<|start_header_id|>user<|end_header_id|>\n", + "post_message": "<|eot_id|>", + }, + "assistant": { + "pre_message": "<|start_header_id|>assistant<|end_header_id|>\n", + "post_message": "<|eot_id|>", + }, }, messages=messages, initial_prompt_value="<|begin_of_text|>", From 0aa8b94ff5e7c1e1cbcf963de492d4c887f6b3ef Mon Sep 17 00:00:00 2001 From: alisalim17 Date: Mon, 29 Apr 2024 18:38:12 +0400 Subject: [PATCH 2/2] test: completion with Cohere command-r-plus model --- litellm/tests/test_completion.py | 70 ++++++++++++++++++++++++++++++++ 1 file changed, 70 insertions(+) diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index fe4aa9c1c8..0174cdaac5 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -231,6 +231,76 @@ def test_completion_claude_3_function_call(): pytest.fail(f"Error occurred: {e}") +def test_completion_cohere_command_r_plus_function_call(): + litellm.set_verbose = True + tools = [ + { + "type": "function", + "function": { + "name": "get_current_weather", + "description": "Get the current weather in a given location", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "The city and state, e.g. San Francisco, CA", + }, + "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}, + }, + "required": ["location"], + }, + }, + } + ] + messages = [ + { + "role": "user", + "content": "What's the weather like in Boston today in Fahrenheit?", + } + ] + try: + # test without max tokens + response = completion( + model="command-r-plus", + messages=messages, + tools=tools, + tool_choice="auto", + ) + # Add any assertions, here to check response args + print(response) + assert isinstance(response.choices[0].message.tool_calls[0].function.name, str) + assert isinstance( + response.choices[0].message.tool_calls[0].function.arguments, str + ) + + messages.append( + response.choices[0].message.model_dump() + ) # Add assistant tool invokes + tool_result = ( + '{"location": "Boston", "temperature": "72", "unit": "fahrenheit"}' + ) + # Add user submitted tool results in the OpenAI format + messages.append( + { + "tool_call_id": response.choices[0].message.tool_calls[0].id, + "role": "tool", + "name": response.choices[0].message.tool_calls[0].function.name, + "content": tool_result, + } + ) + # In the second response, Cohere should deduce answer from tool results + second_response = completion( + model="command-r-plus", + messages=messages, + tools=tools, + tool_choice="auto", + ) + print(second_response) + except Exception as e: + pytest.fail(f"Error occurred: {e}") + + def test_parse_xml_params(): from litellm.llms.prompt_templates.factory import parse_xml_params