fix(token_counter): bound concurrent HuggingFace encodes so a burst of large counts cannot exhaust memory

This commit is contained in:
mateo-berri 2026-09-08 17:38:01 -07:00
parent d202885f8b
commit c6a0381595
3 changed files with 37 additions and 1 deletions

View file

@ -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))

View file

@ -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":

View file

@ -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]