mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(cost): pick each batch price component's tier from its own keys
The batch rate picker crossed one threshold for every component, so a deployment declaring only an output tier also moved its input, cached, and cache-write rates to that cutoff. Each component now crosses its own *_above_<N>k_tokens_batches keys and falls back to its flat key. The JSON schema is regenerated with the generator as it is on main: cost-map-guard renders the PR's cost map with the base branch's generator, so the descriptions for the new batch cache keys move to a follow-up.
This commit is contained in:
parent
1631d2cef8
commit
9d59fb4071
4 changed files with 55 additions and 36 deletions
|
|
@ -222,10 +222,6 @@ COST_DESCRIPTIONS: dict[str, str] = {
|
|||
"cache_read_input_token_cost": "USD per prompt token served from the provider's prompt cache.",
|
||||
"input_cost_per_token_batches": "USD per prompt token via the provider's batch API.",
|
||||
"output_cost_per_token_batches": "USD per generated token via the provider's batch API.",
|
||||
"cache_read_input_token_cost_batches": "USD per cached prompt token via the provider's batch API.",
|
||||
"cache_creation_input_token_cost_batches": (
|
||||
"USD per token written to the provider's prompt cache via its batch API."
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -236,8 +232,6 @@ def cost_description(key: str) -> Optional[str]:
|
|||
return "Flex service-tier rate for the same-named base field."
|
||||
if key.endswith("_priority"):
|
||||
return "Priority service-tier rate for the same-named base field."
|
||||
if "_above_" in key and key.endswith("_batches"):
|
||||
return "Batch API rate applied once the prompt exceeds the token threshold in the field name."
|
||||
if "_above_" in key:
|
||||
return "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -74,7 +74,7 @@ _BATCH_RATE_PREFIXES: Final = (
|
|||
"cache_read_input_token_cost",
|
||||
"cache_creation_input_token_cost",
|
||||
)
|
||||
_BATCH_TIER_KEY: Final = re.compile(rf"^(?:{'|'.join(_BATCH_RATE_PREFIXES)})_above_(\d+k?)_tokens{_BATCH_KEY_SUFFIX}$")
|
||||
_BATCH_TIER_KEY: Final = re.compile(rf"^({'|'.join(_BATCH_RATE_PREFIXES)})_above_(\d+k?)_tokens{_BATCH_KEY_SUFFIX}$")
|
||||
_NON_STANDARD_THRESHOLD_SUFFIXES: Final = (*_SERVICE_TIER_SUFFIXES, _BATCH_KEY_SUFFIX)
|
||||
|
||||
|
||||
|
|
@ -294,36 +294,42 @@ def _batch_tier_rate(model_info: ModelInfo, tier_key: str, flat_key: str) -> flo
|
|||
return _batch_rate(model_info, flat_key) if tier_rate is None else tier_rate
|
||||
|
||||
|
||||
def _batch_rate_for_threshold(model_info: ModelInfo, prefix: str, threshold: str | None) -> float | None:
|
||||
def _batch_tier_thresholds(model_info: ModelInfo, prefix: str) -> frozenset[str]:
|
||||
return frozenset(
|
||||
tier.group(2)
|
||||
for key, value in model_info.items()
|
||||
if value is not None and (tier := _BATCH_TIER_KEY.match(key)) is not None and tier.group(1) == prefix
|
||||
)
|
||||
|
||||
|
||||
def _crossed_batch_tier(model_info: ModelInfo, prefix: str, usage: Usage, inclusive: bool) -> str | None:
|
||||
return next(
|
||||
(
|
||||
threshold
|
||||
for threshold in sorted(
|
||||
_batch_tier_thresholds(model_info, prefix), key=_parse_token_threshold, reverse=True
|
||||
)
|
||||
if _prompt_exceeds_threshold(usage.prompt_tokens, _parse_token_threshold(threshold), inclusive)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def _batch_rate_for_prefix(model_info: ModelInfo, prefix: str, usage: Usage, inclusive: bool) -> float | None:
|
||||
flat_key: Final = f"{prefix}{_BATCH_KEY_SUFFIX}"
|
||||
threshold: Final = _crossed_batch_tier(model_info, prefix, usage, inclusive)
|
||||
if threshold is None:
|
||||
return _batch_rate(model_info, flat_key)
|
||||
return _batch_tier_rate(model_info, f"{prefix}_above_{threshold}_tokens{_BATCH_KEY_SUFFIX}", flat_key)
|
||||
|
||||
|
||||
def _batch_tier_thresholds(model_info: ModelInfo) -> frozenset[str]:
|
||||
return frozenset(
|
||||
tier.group(1)
|
||||
for key, value in model_info.items()
|
||||
if value is not None and (tier := _BATCH_TIER_KEY.match(key)) is not None
|
||||
)
|
||||
|
||||
|
||||
def get_batch_cost_rates(model_info: ModelInfo, usage: Usage, custom_llm_provider: str | None) -> BatchCostRates:
|
||||
inclusive: Final = _uses_inclusive_token_thresholds(custom_llm_provider)
|
||||
crossed_threshold: Final = next(
|
||||
(
|
||||
threshold
|
||||
for threshold in sorted(_batch_tier_thresholds(model_info), key=_parse_token_threshold, reverse=True)
|
||||
if _prompt_exceeds_threshold(usage.prompt_tokens, _parse_token_threshold(threshold), inclusive)
|
||||
),
|
||||
None,
|
||||
)
|
||||
return BatchCostRates(
|
||||
input=_batch_rate_for_threshold(model_info, "input_cost_per_token", crossed_threshold),
|
||||
output=_batch_rate_for_threshold(model_info, "output_cost_per_token", crossed_threshold),
|
||||
cache_read=_batch_rate_for_threshold(model_info, "cache_read_input_token_cost", crossed_threshold),
|
||||
cache_creation=_batch_rate_for_threshold(model_info, "cache_creation_input_token_cost", crossed_threshold),
|
||||
input=_batch_rate_for_prefix(model_info, "input_cost_per_token", usage, inclusive),
|
||||
output=_batch_rate_for_prefix(model_info, "output_cost_per_token", usage, inclusive),
|
||||
cache_read=_batch_rate_for_prefix(model_info, "cache_read_input_token_cost", usage, inclusive),
|
||||
cache_creation=_batch_rate_for_prefix(model_info, "cache_creation_input_token_cost", usage, inclusive),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -112,7 +112,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": {
|
||||
"type": "number",
|
||||
|
|
@ -126,8 +126,7 @@
|
|||
},
|
||||
"cache_creation_input_token_cost_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "USD per token written to the provider's prompt cache via its batch API."
|
||||
"minimum": 0
|
||||
},
|
||||
"cache_creation_input_token_cost_flex": {
|
||||
"type": "number",
|
||||
|
|
@ -176,7 +175,7 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": {
|
||||
"type": "number",
|
||||
|
|
@ -195,8 +194,7 @@
|
|||
},
|
||||
"cache_read_input_token_cost_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "USD per cached prompt token via the provider's batch API."
|
||||
"minimum": 0
|
||||
},
|
||||
"cache_read_input_token_cost_flex": {
|
||||
"type": "number",
|
||||
|
|
@ -353,7 +351,7 @@
|
|||
"input_cost_per_token_above_272k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_token_above_272k_tokens_flex": {
|
||||
"type": "number",
|
||||
|
|
@ -557,7 +555,7 @@
|
|||
"output_cost_per_token_above_272k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"output_cost_per_token_above_272k_tokens_flex": {
|
||||
"type": "number",
|
||||
|
|
|
|||
|
|
@ -3743,6 +3743,27 @@ def test_get_batch_cost_rates_crosses_a_tier_declared_without_an_input_tier_key(
|
|||
assert rates.input == 1e-6
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("prompt_tokens", "expected_input", "expected_output"),
|
||||
[(200_000, 1e-6, 4e-6), (250_000, 1e-6, 5e-6), (300_000, 2e-6, 5e-6)],
|
||||
)
|
||||
def test_get_batch_cost_rates_crosses_each_components_own_tier(prompt_tokens, expected_input, expected_output):
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import get_batch_cost_rates
|
||||
|
||||
rates = get_batch_cost_rates(
|
||||
_batch_rates_model_info(
|
||||
input_cost_per_token_batches=1e-6,
|
||||
input_cost_per_token_above_272k_tokens_batches=2e-6,
|
||||
output_cost_per_token_batches=4e-6,
|
||||
output_cost_per_token_above_200k_tokens_batches=5e-6,
|
||||
),
|
||||
Usage(prompt_tokens=prompt_tokens, completion_tokens=1, total_tokens=prompt_tokens + 1),
|
||||
"openai",
|
||||
)
|
||||
|
||||
assert (rates.input, rates.output) == (expected_input, expected_output)
|
||||
|
||||
|
||||
def test_get_batch_cost_rates_has_no_cache_write_rate_without_a_cache_write_batch_key():
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import get_batch_cost_rates
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue