From fef2a3984339119f79f83374d4db0083d2849486 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 13 Sep 2023 19:39:11 -0700 Subject: [PATCH] adding finish reason mapping for aleph alpha and baseten --- litellm/llms/aleph_alpha.py | 1 + litellm/llms/baseten.py | 1 + litellm/tests/test_completion.py | 57 +++++++++++++++++--------------- pyproject.toml | 2 +- 4 files changed, 33 insertions(+), 28 deletions(-) diff --git a/litellm/llms/aleph_alpha.py b/litellm/llms/aleph_alpha.py index c5e8cded2dc..899fa3cdc1b 100644 --- a/litellm/llms/aleph_alpha.py +++ b/litellm/llms/aleph_alpha.py @@ -100,6 +100,7 @@ def completion( else: try: model_response["choices"][0]["message"]["content"] = completion_response["completions"][0]["completion"] + model_response.choices[0].finish_reason = completion_response["completions"][0]["finish_reason"] except: raise AlephAlphaError(message=json.dumps(completion_response), status_code=response.status_code) diff --git a/litellm/llms/baseten.py b/litellm/llms/baseten.py index 5ed39416190..8f24e129ea4 100644 --- a/litellm/llms/baseten.py +++ b/litellm/llms/baseten.py @@ -117,6 +117,7 @@ def completion( model_response["choices"][0]["message"]["content"] = completion_response[0]["generated_text"] ## GETTING LOGPROBS if "details" in completion_response[0] and "tokens" in completion_response[0]["details"]: + model_response.choices[0].finish_reason = completion_response[0]["details"]["finish_reason"] sum_logprob = 0 for token in completion_response[0]["details"]["tokens"]: sum_logprob += token["logprob"] diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index 934354c2c14..9f3e9ed3ce3 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -487,25 +487,25 @@ def test_completion_azure_deployment_id(): # Replicate API endpoints are unstable -> throw random CUDA errors -> this means our tests can fail even if our tests weren't incorrect. -def test_completion_replicate_llama_2(): - model_name = "replicate/llama-2-70b-chat:2796ee9483c3fd7aa2e171d38f4ca12251a30609463dcfd4cd76703f22e96cdf" - try: - response = completion( - model=model_name, - messages=messages, - max_tokens=20, - custom_llm_provider="replicate" - ) - print(response) - cost = completion_cost(completion_response=response) - print("Cost for completion call with llama-2: ", f"${float(cost):.10f}") - # Add any assertions here to check the response - response_str = response["choices"][0]["message"]["content"] - print(response_str) - if type(response_str) != str: - pytest.fail(f"Error occurred: {e}") - except Exception as e: - pytest.fail(f"Error occurred: {e}") +# def test_completion_replicate_llama_2(): +# model_name = "replicate/llama-2-70b-chat:2796ee9483c3fd7aa2e171d38f4ca12251a30609463dcfd4cd76703f22e96cdf" +# try: +# response = completion( +# model=model_name, +# messages=messages, +# max_tokens=20, +# custom_llm_provider="replicate" +# ) +# print(response) +# cost = completion_cost(completion_response=response) +# print("Cost for completion call with llama-2: ", f"${float(cost):.10f}") +# # Add any assertions here to check the response +# response_str = response["choices"][0]["message"]["content"] +# print(response_str) +# if type(response_str) != str: +# pytest.fail(f"Error occurred: {e}") +# except Exception as e: +# pytest.fail(f"Error occurred: {e}") # test_completion_replicate_llama_2() def test_completion_replicate_vicuna(): @@ -601,6 +601,7 @@ def test_completion_sagemaker(): except Exception as e: pytest.fail(f"Error occurred: {e}") +# test_completion_sagemaker() ######## Test VLLM ######## # def test_completion_vllm(): # try: @@ -658,14 +659,15 @@ def test_completion_sagemaker(): # test_completion_custom_api_base() -# def test_vertex_ai(): -# model_name = "chat-bison" -# try: -# response = completion(model=model_name, messages=messages, logger_fn=logger_fn) -# print(response) -# except Exception as e: -# pytest.fail(f"Error occurred: {e}") +def test_vertex_ai(): + model_name = "chat-bison" + try: + response = completion(model=model_name, messages=messages, logger_fn=logger_fn) + print(response) + except Exception as e: + pytest.fail(f"Error occurred: {e}") +test_vertex_ai() # def test_petals(): # model_name = "stabilityai/StableBeluga2" # try: @@ -696,12 +698,13 @@ def test_completion_with_fallbacks(): # def test_baseten(): # try: -# response = completion(model="baseten/RqgAEn0", messages=messages, logger_fn=logger_fn) +# response = completion(model="baseten/7qQNLDB", messages=messages, logger_fn=logger_fn) # # Add any assertions here to check the response # print(response) # except Exception as e: # pytest.fail(f"Error occurred: {e}") +# test_baseten() # def test_baseten_falcon_7bcompletion(): # model_name = "qvv0xeq" # try: diff --git a/pyproject.toml b/pyproject.toml index 6585ca5ba78..5b82a3b5f0b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "0.1.621" +version = "0.1.622" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT License"