From 3829ebdce54a0330ca2c9d60a5bbee94c2b09382 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 7 Sep 2026 18:26:23 -0700 Subject: [PATCH] fix(token_counter): estimate over-cap strings from evenly spaced samples instead of the prefix A string above TOKEN_COUNTER_MAX_EXACT_CHARS was counted from its first cap characters and scaled, so a string whose start tokenizes unlike its end got a skewed count, and that count reaches fallback billing when the provider sends no usage. The estimate now tokenizes 16 evenly spaced samples that together total the cap and scales their sum by the string's length, keeping the same bound on work while tracking the whole string --- litellm/litellm_core_utils/token_counter.py | 16 ++++++++- .../litellm_core_utils/test_token_counter.py | 36 +++++++++++-------- 2 files changed, 37 insertions(+), 15 deletions(-) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index d526f0f2996..e3a0d4c785a 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -315,6 +315,8 @@ def calculate_img_tokens( TokenCounterFunction = Callable[[str], int] + +EXTRAPOLATION_SAMPLES: Final = 16 """ Type for a function that counts tokens in a string. """ @@ -547,11 +549,23 @@ def _get_extrapolating_count_function( def count_tokens(text: str) -> int: if len(text) <= max_exact_chars: return count_exactly(text) - return round(count_exactly(text[:max_exact_chars]) * len(text) / max_exact_chars) + samples: Final = _evenly_spaced_samples(text, max_exact_chars) + sampled_chars: Final = sum(len(sample) for sample in samples) + return round(sum(count_exactly(sample) for sample in samples) * len(text) / sampled_chars) return count_tokens +def _evenly_spaced_samples(text: str, total_chars: int) -> tuple[str, ...]: + sample_count: Final = min(EXTRAPOLATION_SAMPLES, total_chars) + sample_chars: Final = total_chars // sample_count + last_start: Final = len(text) - sample_chars + return tuple( + text[start : start + sample_chars] + for start in (last_start * index // max(sample_count - 1, 1) for index in range(sample_count)) + ) + + def _get_count_function( model: str | None, custom_tokenizer: dict | SelectTokenizerResponse | None = None, diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index fddeacc37f9..5da8413ac9b 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -153,27 +153,35 @@ async def test_huggingface_count_in_a_worker_thread_leaves_the_event_loop_free() assert max(lags) < took / 4, f"the event loop stalled {max(lags):.3f}s during a {took:.3f}s count" -@pytest.mark.parametrize( - ("max_exact_chars", "expected"), - [(1_000, 10_000), (5_000, 6_000), (10_000, 6_000)], -) -def test_count_above_the_cap_scales_the_exact_count_of_the_prefix(max_exact_chars, expected): - def count_exactly(chunk: str) -> int: - return chunk.count("a") + len(chunk) +@pytest.mark.parametrize("max_exact_chars", [64, 1_000, 2_500]) +def test_count_above_the_cap_samples_the_whole_string_and_scales(max_exact_chars): + count_exactly: Final = MagicMock(side_effect=lambda chunk: chunk.count("a") + len(chunk)) + front_heavy: Final = "a" * 1_000 + "b" * 4_000 + exact: Final = 1_000 + len(front_heavy) - count_tokens: Final = _get_extrapolating_count_function(count_exactly, max_exact_chars=max_exact_chars) + estimate: Final = _get_extrapolating_count_function(count_exactly, max_exact_chars=max_exact_chars)(front_heavy) - assert count_tokens("a" * 1_000 + "b" * 4_000) == expected + assert abs(estimate - exact) <= exact // 100 + assert sum(len(call.args[0]) for call in count_exactly.call_args_list) <= max_exact_chars + + +def test_count_at_or_below_the_cap_is_exact(): + count_exactly: Final = MagicMock(side_effect=len) + + assert _get_extrapolating_count_function(count_exactly, max_exact_chars=5_000)("a" * 5_000) == 5_000 + assert count_exactly.call_args_list == [(("a" * 5_000,),)] def test_token_counter_applies_the_default_cap(): max_exact_chars: Final = litellm.constants.TOKEN_COUNTER_MAX_EXACT_CHARS - prefix: Final = ("The quick brown fox jumps over the lazy dog. " * (max_exact_chars // 45 + 1))[:max_exact_chars] - over_the_cap: Final = prefix + "a" * 200_000 - scaled: Final = round(token_counter_new(model="gpt-5.6", text=prefix) * len(over_the_cap) / max_exact_chars) + prose: Final = ("The quick brown fox jumps over the lazy dog. " * (max_exact_chars // 45 + 1))[:max_exact_chars] + over_the_cap: Final = prose + "a" * 200_000 + exact: Final = _get_exact_count_function("gpt-5.6")(over_the_cap) - assert token_counter_new(model="gpt-5.6", text=over_the_cap) == scaled - assert _get_exact_count_function("gpt-5.6")(over_the_cap) != scaled + estimate: Final = token_counter_new(model="gpt-5.6", text=over_the_cap) + + assert estimate != exact + assert abs(estimate - exact) <= exact // 100 @pytest.mark.parametrize(