feat(cost_calculator.py): refactor rerank cost tracking to support bedrock cost tracking

This commit is contained in:
Krrish Dholakia 2025-02-20 14:44:24 -08:00
parent 4e7d3523af
commit b86a139b64
5 changed files with 53 additions and 79 deletions

View file

@ -16,15 +16,9 @@ from litellm.llms.anthropic.cost_calculation import (
from litellm.llms.azure.cost_calculation import (
cost_per_token as azure_openai_cost_per_token,
)
from litellm.llms.azure_ai.cost_calculator import (
cost_per_query as azure_ai_rerank_cost_per_query,
)
from litellm.llms.bedrock.image.cost_calculator import (
cost_calculator as bedrock_image_cost_calculator,
)
from litellm.llms.cohere.cost_calculator import (
cost_per_query as cohere_rerank_cost_per_query,
)
from litellm.llms.databricks.cost_calculator import (
cost_per_token as databricks_cost_per_token,
)
@ -54,7 +48,9 @@ from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.rerank import RerankResponse
from litellm.types.utils import (
CallTypesLiteral,
LlmProviders,
LlmProvidersSet,
ModelInfo,
PassthroughCallTypes,
Usage,
)
@ -64,6 +60,7 @@ from litellm.utils import (
EmbeddingResponse,
ImageResponse,
ModelResponse,
ProviderConfigManager,
TextCompletionResponse,
TranscriptionResponse,
_cached_get_model_info_helper,
@ -847,16 +844,22 @@ def rerank_cost(
)
try:
if custom_llm_provider == "cohere":
return cohere_rerank_cost_per_query(
model=model, num_queries=default_num_queries
config = ProviderConfigManager.get_provider_rerank_config(
model=model, provider=LlmProviders(custom_llm_provider)
)
try:
model_info: Optional[ModelInfo] = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
elif custom_llm_provider == "azure_ai":
return azure_ai_rerank_cost_per_query(
model=model, num_queries=default_num_queries
)
raise ValueError(
f"invalid custom_llm_provider for rerank model: {model}, custom_llm_provider: {custom_llm_provider}"
except Exception:
model_info = None
return config.calculate_rerank_cost(
model=model,
custom_llm_provider=custom_llm_provider,
num_queries=default_num_queries,
model_info=model_info,
)
except Exception as e:
raise e

View file

@ -1,32 +0,0 @@
"""
Handles custom cost calculation for Azure AI models.
Custom cost calculation for Azure AI models only requied for rerank.
"""
from typing import Tuple
from litellm.utils import get_model_info
def cost_per_query(model: str, num_queries: int = 1) -> Tuple[float, float]:
"""
Calculates the cost per query for a given rerank model.
Input:
- model: str, the model name without provider prefix
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
model_info = get_model_info(model=model, custom_llm_provider="azure_ai")
if (
"input_cost_per_query" not in model_info
or model_info["input_cost_per_query"] is None
):
return 0.0, 0.0
prompt_cost = model_info["input_cost_per_query"] * num_queries
return prompt_cost, 0.0

View file

@ -1,9 +1,10 @@
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm.types.rerank import OptionalRerankParams, RerankResponse
from litellm.types.utils import ModelInfo
from ..chat.transformation import BaseLLMException
@ -84,3 +85,34 @@ class BaseRerankConfig(ABC):
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
pass
@staticmethod
def calculate_rerank_cost(
model: str,
custom_llm_provider: Optional[str] = None,
num_queries: int = 1,
model_info: Optional[ModelInfo] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per query for a given rerank model.
Input:
- model: str, the model name without provider prefix
- custom_llm_provider: str, the provider used for the model. If provided, used to check if the litellm model info is for that provider.
- num_queries: int, the number of queries to calculate the cost for
- model_info: ModelInfo, the model info for the given model
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
if (
model_info is None
or "input_cost_per_query" not in model_info
or model_info["input_cost_per_query"] is None
):
return 0.0, 0.0
prompt_cost = model_info["input_cost_per_query"] * num_queries
return prompt_cost, 0.0

View file

@ -1,31 +0,0 @@
"""
Custom cost calculator for Cohere rerank models
"""
from typing import Tuple
from litellm.utils import get_model_info
def cost_per_query(model: str, num_queries: int = 1) -> Tuple[float, float]:
"""
Calculates the cost per query for a given rerank model.
Input:
- model: str, the model name without provider prefix
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
model_info = get_model_info(model=model, custom_llm_provider="cohere")
if (
"input_cost_per_query" not in model_info
or model_info["input_cost_per_query"] is None
):
return 0.0, 0.0
prompt_cost = model_info["input_cost_per_query"] * num_queries
return prompt_cost, 0.0

View file

@ -95,6 +95,8 @@ class BaseLLMRerankTest(ABC):
assert response.id is not None
assert response.results is not None
assert response._hidden_params["response_cost"] is not None
assert_response_shape(
response=response, custom_llm_provider=custom_llm_provider.value
)