wkea/blockdiffusion-api
0
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 