From c6a0381595b3085937e514c7af3f54c2ceb77598 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 17:38:01 -0700 Subject: [PATCH] fix(token_counter): bound concurrent HuggingFace encodes so a burst of large counts cannot exhaust memory --- litellm/constants.py | 1 + litellm/litellm_core_utils/token_counter.py | 6 +++- .../litellm_core_utils/test_token_counter.py | 31 +++++++++++++++++++ 3 files changed, 37 insertions(+), 1 deletion(-) diff --git a/litellm/constants.py b/litellm/constants.py index a9112c90a54..05d442dd070 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -397,6 +397,7 @@ TOKEN_COUNTER_MAX_EXACT_CHARS: Final = get_env_int_in_range( minimum=1, maximum=1_000_000_000, ) +TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES: Final = 4 MAX_TILE_WIDTH: Final = int(os.getenv("MAX_TILE_WIDTH", 512)) MAX_TILE_HEIGHT: Final = int(os.getenv("MAX_TILE_HEIGHT", 512)) OPENAI_FILE_SEARCH_COST_PER_1K_CALLS: Final = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000)) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index e3af64fd3bf..8624e79faea 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -3,6 +3,7 @@ import base64 import io import struct +import threading from collections.abc import Callable, Iterable, Mapping, Sequence from typing import Final, Literal, cast @@ -22,6 +23,7 @@ from litellm.constants import ( MAX_TILE_HEIGHT, MAX_TILE_WIDTH, TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS, + TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES, TOKEN_COUNTER_MAX_EXACT_CHARS, ) from litellm.litellm_core_utils.default_encoding import encoding as default_encoding @@ -320,6 +322,7 @@ Type for a function that counts tokens in a string. """ EXTRAPOLATION_SAMPLES: Final = 16 +_HF_ENCODE_SLOTS: Final = threading.BoundedSemaphore(TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES) def _get_tiktoken_count_function( @@ -587,7 +590,8 @@ def _get_exact_count_function( tokenizer: Final[Tokenizer] = tokenizer_json["tokenizer"] def count_tokens(text: str) -> int: - return len(tokenizer.encode_batch_fast([text])[0]) + with _HF_ENCODE_SLOTS: + return len(tokenizer.encode_batch_fast([text])[0]) return count_tokens elif tokenizer_json["type"] == "openai_tokenizer": 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 e3dd5feb18d..655be5146cd 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -2,8 +2,10 @@ # This tests litellm.token_counter.token_counter() function import asyncio import importlib +import threading import time import traceback +from concurrent.futures import ThreadPoolExecutor from typing import Final from unittest.mock import MagicMock @@ -16,6 +18,7 @@ import litellm from litellm import create_pretrained_tokenizer, decode, encode, get_modified_max_tokens from litellm import token_counter as token_counter_old import litellm.constants +from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.token_counter import ( _get_exact_count_function, @@ -162,6 +165,34 @@ def test_count_at_or_below_the_cap_is_exact(): assert count_exactly.call_args_list == [(("a" * 5_000,),)] +class _SlowEncoder: + def __init__(self) -> None: + self._lock: Final = threading.Lock() + self.in_flight = 0 + self.peak_in_flight = 0 + + def encode_batch_fast(self, texts: list[str]) -> list[list[int]]: + with self._lock: + self.in_flight += 1 + self.peak_in_flight = max(self.peak_in_flight, self.in_flight) + time.sleep(0.1) + with self._lock: + self.in_flight -= 1 + return [[0] * len(text) for text in texts] + + +def test_huggingface_counts_run_at_most_the_configured_number_at_once(): + encoder: Final = _SlowEncoder() + count: Final = _get_exact_count_function(None, {"type": "huggingface_tokenizer", "tokenizer": encoder}) + burst: Final = 2 * TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES + + with ThreadPoolExecutor(max_workers=burst) as pool: + counts: Final = tuple(pool.map(count, ["abc"] * burst)) + + assert counts == (3,) * burst + assert 1 < encoder.peak_in_flight <= TOKEN_COUNTER_MAX_CONCURRENT_HF_ENCODES + + def test_token_counter_applies_the_default_cap(): max_exact_chars: Final = litellm.constants.TOKEN_COUNTER_MAX_EXACT_CHARS prose: Final = ("The quick brown fox jumps over the lazy dog. " * (max_exact_chars // 45 + 1))[:max_exact_chars]