Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
simple_tokenizer.py145 linesDownload Raw Back to text
1from __future__ import annotations
2
3import re
4from dataclasses import dataclass
5from typing import Dict, Iterable, List, Optional, Sequence
6
7
8def _is_cjk(ch: str) -> bool:
9    """
10    粗略判断 CJK 字符(覆盖常见中文/日文/韩文范围)。
11    """
12    code = ord(ch)
13    return (
14        0x4E00 <= code <= 0x9FFF  # CJK Unified Ideographs
15        or 0x3400 <= code <= 0x4DBF  # CJK Unified Ideographs Extension A
16        or 0x3040 <= code <= 0x30FF  # Japanese Hiragana/Katakana
17        or 0xAC00 <= code <= 0xD7AF  # Hangul Syllables
18    )
19
20
21def tokenize(text: str) -> List[str]:
22    """
23    轻量 tokenizer:
24    - 英文/数字按“词”切分
25    - 中文(CJK)按“单字”切分(更适合无空格中文提示词)
26    """
27    text = (text or "").strip().lower()
28    if not text:
29        return []
30
31    tokens: List[str] = []
32    buf: List[str] = []
33
34    def flush_buf() -> None:
35        if buf:
36            tokens.append("".join(buf))
37            buf.clear()
38
39    for ch in text:
40        if _is_cjk(ch):
41            flush_buf()
42            tokens.append(ch)
43            continue
44        # 英文/数字累计成 word,其他当作分隔符
45        if ch.isalnum() or ch in {"_", "-"}:
46            buf.append(ch)
47        else:
48            flush_buf()
49    flush_buf()
50    return tokens
51
52
53@dataclass
54class VocabConfig:
55    max_vocab: int = 8000
56    min_freq: int = 1
57
58
59class SimpleVocab:
60    """
61    极简词表:
62    - id=0: <pad>
63    - id=1: <unk>
64    """
65
66    PAD = "<pad>"
67    UNK = "<unk>"
68
69    def __init__(self, token_to_id: Dict[str, int]):
70        if self.PAD not in token_to_id or self.UNK not in token_to_id:
71            raise ValueError("token_to_id 必须包含 <pad>/<unk>")
72        self.token_to_id = token_to_id
73        self.id_to_token = {i: t for t, i in token_to_id.items()}
74
75    @property
76    def pad_id(self) -> int:
77        return int(self.token_to_id[self.PAD])
78
79    @property
80    def unk_id(self) -> int:
81        return int(self.token_to_id[self.UNK])
82
83    @property
84    def size(self) -> int:
85        return len(self.token_to_id)
86
87    @classmethod
88    def build(cls, texts: Iterable[str], cfg: VocabConfig, extra_tokens: Optional[Sequence[str]] = None) -> "SimpleVocab":
89        """
90        构建词表。
91        - texts:训练语料
92        - extra_tokens:可选“强制加入”的 token 文本(会先 tokenize 再加入)
93
94        说明:
95        - 续训时 tokenizer/vocab 必须与 checkpoint 一致,否则 embedding 尺寸会不匹配;
96          因此 extra_tokens 仅适用于“新训练”阶段。
97        """
98        freq: Dict[str, int] = {}
99        for t in texts:
100            for tok in tokenize(t):
101                freq[tok] = freq.get(tok, 0) + 1
102
103        # 按频率排序,截断到 max_vocab(保留 PAD/UNK)
104        items = sorted(freq.items(), key=lambda x: (-x[1], x[0]))
105        items = [it for it in items if it[1] >= cfg.min_freq]
106
107        token_to_id = {cls.PAD: 0, cls.UNK: 1}
108
109        # 先插入强制 token(优先级高于按频率截断)
110        if extra_tokens:
111            forced: List[str] = []
112            for t in extra_tokens:
113                forced.extend(tokenize(t))
114            for tok in forced:
115                if tok in token_to_id:
116                    continue
117                if len(token_to_id) >= int(cfg.max_vocab):
118                    break
119                token_to_id[tok] = len(token_to_id)
120
121        # 再按频率补齐到 max_vocab
122        for tok, _ in items:
123            if tok in token_to_id:
124                continue
125            if len(token_to_id) >= int(cfg.max_vocab):
126                break
127            token_to_id[tok] = len(token_to_id)
128        return cls(token_to_id)
129
130    def encode(self, text: str, max_len: int) -> List[int]:
131        toks = tokenize(text)
132        ids = [self.token_to_id.get(t, self.unk_id) for t in toks[:max_len]]
133        if len(ids) < max_len:
134            ids.extend([self.pad_id] * (max_len - len(ids)))
135        return ids
136
137    def to_dict(self) -> Dict:
138        return {"token_to_id": self.token_to_id}
139
140    @classmethod
141    def from_dict(cls, d: Dict) -> "SimpleVocab":
142        return cls(token_to_id=dict(d["token_to_id"]))
143
144
145