From 22ba5fa186c921457871e5b4322fe0e49a734f8a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 16 May 2024 10:59:29 -0700 Subject: [PATCH] feat - try using hf tokenizer --- litellm/proxy/_types.py | 4 ++++ litellm/proxy/proxy_server.py | 21 +++++++++++++++++++-- litellm/utils.py | 13 +++++++++++-- 3 files changed, 34 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 63065705619..4d89bbd9f26 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -79,6 +79,8 @@ class LiteLLMRoutes(enum.Enum): "/v1/models", ] + llm_utils_routes: List = ["utils/token_counter"] + info_routes: List = [ "/key/info", "/team/info", @@ -1012,3 +1014,5 @@ class TokenCountRequest(LiteLLMBase): class TokenCountResponse(LiteLLMBase): total_tokens: int model: str + base_model: str + tokenizer_type: str diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index beb206e179e..690a0cb0b6b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4775,21 +4775,38 @@ async def token_counter(request: TokenCountRequest): """ """ from litellm import token_counter + global llm_router + prompt = request.prompt messages = request.messages + if llm_router is not None: + # get 1 deployment corresponding to the model + for _model in llm_router.model_list: + if _model["model_name"] == request.model: + deployment = _model + break + + litellm_model_name = deployment.get("litellm_params", {}).get("model") + # remove the custom_llm_provider_prefix in the litellm_model_name + if "/" in litellm_model_name: + litellm_model_name = litellm_model_name.split("/", 1)[1] + if prompt is None and messages is None: raise HTTPException( status_code=400, detail="prompt or messages must be provided" ) - total_tokens = token_counter( - model=request.model, + total_tokens, tokenizer_used = token_counter( + model=litellm_model_name, text=prompt, messages=messages, + return_tokenizer_used=True, ) return TokenCountResponse( total_tokens=total_tokens, model=request.model, + base_model=litellm_model_name, + tokenizer_type=tokenizer_used, ) diff --git a/litellm/utils.py b/litellm/utils.py index 36f4ad481f9..1725969b25d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3860,7 +3860,12 @@ def _select_tokenizer(model: str): return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} # default - tiktoken else: - return {"type": "openai_tokenizer", "tokenizer": encoding} + tokenizer = None + try: + tokenizer = Tokenizer.from_pretrained(model) + return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} + except: + return {"type": "openai_tokenizer", "tokenizer": encoding} def encode(model="", text="", custom_tokenizer: Optional[dict] = None): @@ -4097,6 +4102,7 @@ def token_counter( text: Optional[Union[str, List[str]]] = None, messages: Optional[List] = None, count_response_tokens: Optional[bool] = False, + return_tokenizer_used: Optional[bool] = False, ): """ Count the number of tokens in a given text using a specified model. @@ -4189,7 +4195,10 @@ def token_counter( ) else: num_tokens = len(encoding.encode(text, disallowed_special=())) # type: ignore - + _tokenizer_type = tokenizer_json["type"] + if return_tokenizer_used: + # used by litellm proxy server -> POST /utils/token_counter + return num_tokens, _tokenizer_type return num_tokens