fix(litellm_cost_calc/google.py): support meta llama vertex ai cost tracking

This commit is contained in:
Krrish Dholakia 2024-07-25 22:11:32 -07:00
parent 2626cc6d30
commit 2f773d9cb6
5 changed files with 25 additions and 11 deletions

View file

@ -44,7 +44,7 @@ def cost_router(
Returns
- str, the specific google cost calc function it should route to.
"""
if custom_llm_provider == "vertex_ai" and "claude" in model:
if custom_llm_provider == "vertex_ai" and ("claude" in model or "llama" in model):
return "cost_per_token"
elif custom_llm_provider == "gemini":
return "cost_per_token"

View file

@ -1,11 +1,4 @@
model_list:
- model_name: "test-model"
- model_name: "gpt-3.5-turbo"
litellm_params:
model: "openai/text-embedding-ada-002"
- model_name: "my-custom-model"
litellm_params:
model: "my-custom-llm/my-model"
litellm_settings:
custom_provider_map:
- {"provider": "my-custom-llm", "custom_handler": custom_handler.my_custom_llm}
model: "openai/gpt-3.5-turbo"

View file

@ -901,7 +901,12 @@ from litellm.tests.test_completion import response_format_tests
@pytest.mark.parametrize(
"model", ["vertex_ai/meta/llama3-405b-instruct-maas"]
) # "vertex_ai",
@pytest.mark.parametrize("sync_mode", [True, False]) # "vertex_ai",
@pytest.mark.parametrize(
"sync_mode",
[
True,
],
) # False
@pytest.mark.asyncio
async def test_llama_3_httpx(model, sync_mode):
try:
@ -932,6 +937,8 @@ async def test_llama_3_httpx(model, sync_mode):
response_format_tests(response=response)
print(f"response: {response}")
assert False
except litellm.RateLimitError as e:
pass
except Exception as e:

View file

@ -907,6 +907,17 @@ def test_vertex_ai_gemini_predict_cost():
assert predictive_cost > 0
def test_vertex_ai_llama_predict_cost():
model = "meta/llama3-405b-instruct-maas"
messages = [{"role": "user", "content": "Hey, hows it going???"}]
custom_llm_provider = "vertex_ai"
predictive_cost = completion_cost(
model=model, messages=messages, custom_llm_provider=custom_llm_provider
)
assert predictive_cost == 0
@pytest.mark.parametrize("model", ["openai/tts-1", "azure/tts-1"])
def test_completion_cost_tts(model):
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"

View file

@ -4919,6 +4919,9 @@ def get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Mod
azure_llms = litellm.azure_llms
if model in azure_llms:
model = azure_llms[model]
if custom_llm_provider is not None and custom_llm_provider == "vertex_ai":
if "meta/" + model in litellm.vertex_llama3_models:
model = "meta/" + model
##########################
if custom_llm_provider is None:
# Get custom_llm_provider