import os from dataclasses import dataclass from typing import Any, Dict, List, Optional, Union from tokenizers import Regex, Tokenizer from tokenizers.decoders import ByteLevel as ByteLevelDecoder from tokenizers.models import BPE from tokenizers.normalizers import NFKC, Sequence as NormalizerSequence from tokenizers.pre_tokenizers import ( ByteLevel, Digits, Sequence as PreTokenizerSequence, Split, ) from tokenizers.trainers import BpeTrainer from transformers import PreTrainedTokenizerFast SPECIAL_TOKENS = ["", "", "", "", "[EOD]", "<|eod|>"] EMOJIS = [ "\U0001F602", "\U0001F62D", "\u2728", "\U0001F680", "\U0001F44D", "\U0001F64F", "\U0001F525", "\U0001F60A", "\u2764\ufe0f", "\U0001F914", "\U0001F923", "\U0001F60D", "\U0001F480", "\U0001F4AF", "\u26a0\ufe0f", "\u2705", "\u274c", "\U0001F4CA", "\U0001F4BB", "\U0001F4F1", "\U0001F623", "\U0001F970", "\U0001F605", "\U0001F606", "\U0001F979", "\U0001F61A", "\U0001F917", "\U0001F61D", "\U0001F440", ] EMOTICONS = [ ":-)", ":)", ":D", ":(", ";)", "XD", "OwO", "UwU", "T_T", "QAQ", "¯\\_(ツ)_/¯", ] MATH_LATEX = [ "\\alpha", "\\beta", "\\gamma", "\\theta", "\\pi", "\\sigma", "\\omega", "\\sum", "\\int", "\\approx", "\\neq", "\\le", "\\ge", "\\infty", "\\partial", "\\nabla", "\\forall", "\\exists", "\\in", "\\notin", "\\rightarrow", "\\Rightarrow", "\\Leftrightarrow", ] CODE_OPERATORS = [ "==", "!=", "<=", ">=", "+=", "-=", "*=", "/=", "=>", "->", "&&", "||", "async", "await", "lambda", ] CUSTOM_TOKENS = [ "<|im_start|>", "<|im_end|>", "<|system|>", "<|user|>", "<|assistant|>", "", "", ] + (EMOJIS + EMOTICONS + MATH_LATEX + CODE_OPERATORS) ALL_TOKENS = SPECIAL_TOKENS + CUSTOM_TOKENS LLM_SPLIT_REGEX = ( r"""(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}|""" r""" ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+""" ) def build_xonelm_tokenizer( corpus: Optional[Union[str, List[str]]] = None, vocab_size: int = 32000, save_path: Optional[str] = None, ) -> PreTrainedTokenizerFast: bpe_model = BPE(unk_token="") tokenizer_raw = Tokenizer(bpe_model) tokenizer_raw.normalizer = NormalizerSequence([NFKC()]) tokenizer_raw.pre_tokenizer = PreTokenizerSequence([ Split(pattern=Regex(LLM_SPLIT_REGEX), behavior="isolated", invert=False), Digits(individual_digits=True), ByteLevel(add_prefix_space=False, use_regex=False), ]) tokenizer_raw.decoder = ByteLevelDecoder() trainer = BpeTrainer( vocab_size=vocab_size, special_tokens=ALL_TOKENS, initial_alphabet=ByteLevel.alphabet(), show_progress=False, ) if corpus is not None: if isinstance(corpus, str) and os.path.isfile(corpus): tokenizer_raw.train([corpus], trainer) elif isinstance(corpus, list) and len(corpus) > 0 and os.path.isfile(corpus[0]): tokenizer_raw.train(corpus, trainer) else: iterator = [corpus] if isinstance(corpus, str) else corpus tokenizer_raw.train_from_iterator(iterator, trainer) else: tokenizer_raw.train_from_iterator(["Hello world 123 \\alpha \\beta == async await"], trainer) if save_path is not None: tokenizer_raw.save(save_path) hf_tokenizer = PreTrainedTokenizerFast( tokenizer_object=tokenizer_raw, bos_token="", eos_token="", pad_token="", unk_token="", additional_special_tokens=ALL_TOKENS, ) return hf_tokenizer @dataclass class SpecialTokenConfig: pad_token_id: int = 0 bos_token_id: int = 1 eos_token_id: int = 2 unk_token_id: int = 3 eod_token_id: int = 4 im_start_id: Optional[int] = None im_end_id: Optional[int] = None separator_token_id: Optional[int] = None class MultiTurnConversationFormatter: def __init__(self, tokenizer: Any, token_config: Optional[SpecialTokenConfig] = None): self.tokenizer = tokenizer self.config = token_config or SpecialTokenConfig() def _get_id(token_str: str) -> Optional[int]: if hasattr(tokenizer, "token_to_id"): return tokenizer.token_to_id(token_str) elif hasattr(tokenizer, "convert_tokens_to_ids"): res = tokenizer.convert_tokens_to_ids(token_str) return res if isinstance(res, int) and res >= 0 else None return None if self.config.im_start_id is None: self.config.im_start_id = _get_id("<|im_start|>") if self.config.im_end_id is None: self.config.im_end_id = _get_id("<|im_end|>") if self.config.eod_token_id is None: self.config.eod_token_id = _get_id("[EOD]") def format_conversation( self, messages: List[Dict[str, str]], max_len: Optional[int] = None ) -> Dict[str, List[int]]: input_ids = [] labels = [] def _encode_text(t: str) -> List[int]: if hasattr(self.tokenizer, "encode"): res = self.tokenizer.encode(t) return res.ids if hasattr(res, "ids") else res elif callable(self.tokenizer): return self.tokenizer(t)["input_ids"] return [] for msg in messages: role = msg["role"] content = msg["content"].strip() header_text = f"<|im_start|>{role}\n" body_text = f"{content}<|im_end|>\n" header_ids = _encode_text(header_text) body_ids = _encode_text(body_text) turn_input_ids = header_ids + body_ids input_ids.extend(turn_input_ids) if role == "assistant": turn_labels = [-100] * len(header_ids) + body_ids labels.extend(turn_labels) else: labels.extend([-100] * len(turn_input_ids)) if self.config.eod_token_id is not None: input_ids.append(self.config.eod_token_id) labels.append(self.config.eod_token_id) if max_len is not None: input_ids = input_ids[:max_len] labels = labels[:max_len] return {"input_ids": input_ids, "labels": labels}