| 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 = ["<s>", "<pad>", "</s>", "<unk>", "[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|>",
|
| "<think>",
|
| "</think>",
|
| ] + (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="<unk>")
|
| 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="<s>",
|
| eos_token="</s>",
|
| pad_token="<pad>",
|
| unk_token="<unk>",
|
| 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} |