compliantLLM / tokenization_compliant_llm.py
Martin Navrátil
Upload 11 files
4689b4d verified
Raw
History Blame Contribute Delete
1.95 kB
"""Exact 256-entry byte-level tokenizer for compliantLLM."""
import json
import os
from transformers import PreTrainedTokenizer
def _byte_token(value):
return f"<0x{value:02X}>"
class CompliantLLMTokenizer(PreTrainedTokenizer):
"""The zero-merge form of byte-level BPE."""
model_input_names = ["input_ids", "attention_mask"]
vocab_files_names = {"vocab_file": "vocab.json"}
def __init__(self, vocab_file=None, **kwargs):
del vocab_file
self.encoder = {_byte_token(value): value for value in range(256)}
self.decoder = {value: token for token, value in self.encoder.items()}
kwargs.setdefault("pad_token", _byte_token(0))
kwargs.setdefault("model_max_length", 1024)
kwargs.setdefault("padding_side", "right")
kwargs.setdefault("truncation_side", "left")
super().__init__(**kwargs)
@property
def vocab_size(self):
return 256
def get_vocab(self):
return dict(self.encoder)
def _tokenize(self, text, **kwargs):
del kwargs
return [_byte_token(value) for value in text.encode("utf-8")]
def _convert_token_to_id(self, token):
return self.encoder.get(token, 0)
def _convert_id_to_token(self, index):
return self.decoder.get(index, _byte_token(0))
def convert_tokens_to_string(self, tokens):
values = [self.encoder[token] for token in tokens if token in self.encoder]
return bytes(values).decode("utf-8", errors="replace")
def save_vocabulary(self, save_directory, filename_prefix=None):
os.makedirs(save_directory, exist_ok=True)
filename = ((filename_prefix + "-") if filename_prefix else "") + "vocab.json"
path = os.path.join(save_directory, filename)
with open(path, "w", encoding="utf-8") as handle:
json.dump(self.encoder, handle, indent=2, sort_keys=True)
handle.write("\n")
return (path,)