XoneLM-1.0-Paper / tokenizer.py
cloverxion's picture
feat: publish XoneLM architecture and LuminaV optimizer
bbf30b7 verified
Raw
History Blame Contribute Delete
6.46 kB
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}