File size: 2,684 Bytes
7b4c0bf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
# Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87.
# Licensed under Apache-2.0; see LICENSE.
import base64
import shutil
import unicodedata
from pathlib import Path

import tiktoken
from transformers import PreTrainedTokenizer


class YuE2Tokenizer(PreTrainedTokenizer):
    vocab_files_names = {"vocab_file": "qwen.tiktoken"}
    model_input_names = ["input_ids"]

    def __init__(self, vocab_file, expected_vocab_size=151643, **kwargs):
        self.vocab_file = str(vocab_file)
        self.ranks = {
            base64.b64decode(token): int(rank)
            for token, rank in (line.split() for line in Path(vocab_file).read_bytes().splitlines() if line)
        }
        if len(self.ranks) != expected_vocab_size:
            raise ValueError(f"Expected {expected_vocab_size} ordinary tokens")
        specials = ["<|endoftext|>", "<|im_start|>", "<|im_end|>", "<R>", "<S>", "<X>", "<mask>", "<sep>"]
        specials += [f"<extra_{i}>" for i in range(200)]
        specials[204:206] = ["<abc>", "</abc>"]
        pattern = r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"
        self.encoding = tiktoken.Encoding(
            "YuE2",
            pat_str=pattern,
            mergeable_ranks=self.ranks,
            special_tokens={s: i + len(self.ranks) for i, s in enumerate(specials)},
        )
        super().__init__(expected_vocab_size=expected_vocab_size, **kwargs)

    @property
    def vocab_size(self):
        return len(self.ranks)

    def get_vocab(self):
        return {base64.b64encode(token).decode(): rank for token, rank in self.ranks.items()}

    def _tokenize(self, text):
        return [base64.b64encode(self.encoding.decode_single_token_bytes(i)).decode() for i in self.encode(text)]

    def _convert_token_to_id(self, token):
        return self.ranks[base64.b64decode(token)]

    def _convert_id_to_token(self, index):
        return base64.b64encode(self.encoding.decode_single_token_bytes(index)).decode()

    def encode(self, text, **kwargs):
        return self.encoding.encode_ordinary(unicodedata.normalize("NFC", text))

    def decode(self, token_ids, **kwargs):
        return self.encoding.decode([int(i) for i in token_ids if 0 <= i < self.encoding.n_vocab], errors="replace")

    def save_vocabulary(self, save_directory, filename_prefix=None):
        target = Path(save_directory) / ((filename_prefix + "-" if filename_prefix else "") + "qwen.tiktoken")
        if target.resolve() != Path(self.vocab_file).resolve():
            shutil.copyfile(self.vocab_file, target)
        return (str(target),)