mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix: preserve tokenizer decode round trips
This commit is contained in:
parent
0c3b4a06cf
commit
2d36596e8b
1 changed files with 22 additions and 0 deletions
|
|
@ -2245,10 +2245,32 @@ def encode(model="", text="", custom_tokenizer: Optional[dict] = None):
|
|||
|
||||
def decode(model="", tokens: List[int] = [], custom_tokenizer: Optional[dict] = None):
|
||||
tokenizer_json = custom_tokenizer or _select_tokenizer(model=model)
|
||||
if tokenizer_json["type"] == "huggingface_tokenizer":
|
||||
tokens = _strip_huggingface_special_token_ids(
|
||||
tokenizer_json["tokenizer"], tokens
|
||||
)
|
||||
dec = tokenizer_json["tokenizer"].decode(tokens)
|
||||
return dec
|
||||
|
||||
|
||||
def _strip_huggingface_special_token_ids(
|
||||
tokenizer: Tokenizer, tokens: List[int]
|
||||
) -> List[int]:
|
||||
try:
|
||||
added_tokens_decoder = tokenizer.get_added_tokens_decoder()
|
||||
except Exception:
|
||||
return tokens
|
||||
|
||||
special_token_ids = {
|
||||
token_id
|
||||
for token_id, added_token in added_tokens_decoder.items()
|
||||
if getattr(added_token, "special", False)
|
||||
}
|
||||
if not special_token_ids:
|
||||
return tokens
|
||||
return [token for token in tokens if token not in special_token_ids]
|
||||
|
||||
|
||||
def create_pretrained_tokenizer(
|
||||
identifier: str, revision="main", auth_token: Optional[str] = None
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue