fix(main.py): fix pydantic warning for usage dict

This commit is contained in:
Krrish Dholakia 2023-12-02 20:02:55 -08:00
parent f0d8a87c48
commit 6c0eec4ff4
2 changed files with 29 additions and 32 deletions

View file

@ -30,7 +30,8 @@ from litellm.utils import (
get_api_key,
mock_completion_streaming_obj,
convert_to_model_response_object,
token_counter
token_counter,
Usage
)
from .llms import (
anthropic,
@ -1288,11 +1289,7 @@ def completion(
model_response["model"] = "ollama/" + model
prompt_tokens = len(encoding.encode(prompt)) # type: ignore
completion_tokens = len(encoding.encode(response_string))
model_response["usage"] = {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
model_response["usage"] = Usage(prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens)
response = model_response
elif (
custom_llm_provider == "baseten"

View file

@ -1,35 +1,35 @@
# ##### THESE TESTS CAN ONLY RUN LOCALLY WITH THE OLLAMA SERVER RUNNING ######
# # https://ollama.ai/
##### THESE TESTS CAN ONLY RUN LOCALLY WITH THE OLLAMA SERVER RUNNING ######
# https://ollama.ai/
# import sys, os
# import traceback
# from dotenv import load_dotenv
# load_dotenv()
# import os
# sys.path.insert(0, os.path.abspath('../..')) # Adds the parent directory to the system path
# import pytest
# import litellm
# from litellm import embedding, completion
# import asyncio
import sys, os
import traceback
from dotenv import load_dotenv
load_dotenv()
import os
sys.path.insert(0, os.path.abspath('../..')) # Adds the parent directory to the system path
import pytest
import litellm
from litellm import embedding, completion
import asyncio
# user_message = "respond in 20 words. who are you?"
# messages = [{ "content": user_message,"role": "user"}]
user_message = "respond in 20 words. who are you?"
messages = [{ "content": user_message,"role": "user"}]
# def test_completion_ollama():
# try:
# response = completion(
# model="ollama/llama2",
# messages=messages,
# max_tokens=200,
# request_timeout = 10,
def test_completion_ollama():
try:
response = completion(
model="ollama/llama2",
messages=messages,
max_tokens=200,
request_timeout = 10,
# )
# print(response)
# except Exception as e:
# pytest.fail(f"Error occurred: {e}")
)
print(response)
except Exception as e:
pytest.fail(f"Error occurred: {e}")
# test_completion_ollama()
test_completion_ollama()
# def test_completion_ollama_with_api_base():
# try: