mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(cost_calculator.py): refactor rerank cost tracking to support bedrock cost tracking
This commit is contained in:
parent
4e7d3523af
commit
b86a139b64
5 changed files with 53 additions and 79 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue