diff --git a/tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py b/tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py index f0f2b8e7f5d..d1a5f78a859 100644 --- a/tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py +++ b/tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py @@ -1,17 +1,37 @@ +from tokenizers import AddedToken, Tokenizer +from tokenizers.models import WordLevel +from tokenizers.pre_tokenizers import Whitespace +from tokenizers.processors import TemplateProcessing + from litellm import decode, encode -def test_decode_can_preserve_huggingface_special_tokens(): - sample_text = "Hello World, this is my input string!" - tokens = encode(model="meta-llama/Llama-2-7b-chat", text=sample_text) +def _create_custom_tokenizer(): + tokenizer = Tokenizer( + WordLevel({"[UNK]": 0, "Hello": 1, "World": 2}, unk_token="[UNK]") + ) + tokenizer.pre_tokenizer = Whitespace() + tokenizer.add_special_tokens([AddedToken("[BOS]", special=True)]) + bos_token_id = tokenizer.token_to_id("[BOS]") + assert bos_token_id is not None + tokenizer.post_processor = TemplateProcessing( + single="[BOS] $A", + special_tokens=[("[BOS]", bos_token_id)], + ) + return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - decoded_text = decode(model="meta-llama/Llama-2-7b-chat", tokens=tokens) + +def test_decode_can_preserve_huggingface_special_tokens(): + custom_tokenizer = _create_custom_tokenizer() + sample_text = "Hello World" + tokens = encode(text=sample_text, custom_tokenizer=custom_tokenizer) + + decoded_text = decode(tokens=tokens, custom_tokenizer=custom_tokenizer) decoded_text_with_special_tokens = decode( - model="meta-llama/Llama-2-7b-chat", tokens=tokens, + custom_tokenizer=custom_tokenizer, skip_special_tokens=False, ) assert decoded_text == sample_text - assert sample_text in decoded_text_with_special_tokens - assert decoded_text_with_special_tokens != sample_text + assert decoded_text_with_special_tokens == "[BOS] Hello World"