mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-26 20:33:15 +00:00
perf(chineseclip): 分词器热路径分配优化 + 差分 oracle 验收(含一处真实语义修复)
嵌入式模型推理(ONNX)本身已是 C++,**可优化的 Go 侧是分词器与预处理**。 本轮先测出成本分布,再改,且**不假设 C 更快**。 ## 实测(合成词表,无需 CHINESECLIP_MODEL_DIR) | 场景 | ns/op | allocs | |---|---:|---:| | short_zh | 25681 | 60 | | short_en | 19100 | 43 | | mid_en | 117343 | 451 | | mid_zh | 189260 | 1047 | | punct_heavy | 283265 | 1388 | | long_zh | 757746 | 4233 | ## pprof 指出的分配源(alloc_objects,mid_zh) - splitOnPunctuation **33.5%**(每 token 都做 []rune + string(cur)) - wordpiece **32.3%**(内层每轮候选都 string(runes[a:b]),多数未命中) - stripAccents/NFD 33.6% cum - basicTokenize 自身只 1.4% ## 改了三处 1. `basicTokenize`:`len([]rune(token))` → `utf8.RuneCountInString` (原为「数个长度」就把整个 token 转 rune 切片) 2. `splitOnPunctuation`:去掉整串 []rune,改逐 rune 扫描 + 一次 flush 3. `wordpiece`:预建 rune 边界表,按字节区间取 substring, 消除「每轮候选都构造 string」 ## ★★ 差分 oracle 抓到一处**真实语义缺陷**(非测量噪声) 本机无模型产物,权威的 TestTokenizerMatchesOfficialReference 会 **SKIP** ⇒ 仅靠现有测试,我的重写**没有被有效验证**。故把改动前的实现原样内联为 oracle 做差分(split/wordpiece/basicTokenize/Encode 四组 + 随机字节 2 万组 + 随机 rune 5000 组)。 它立刻抓到:`"\xbc\xef=..."` 旧实现得 `["��" ...]`,新实现得 `["\xbc\xef" ...]`。 根因是 `[]rune(s)` 会把**非法字节归一成 U+FFFD**,而纯字节切片原样保留坏字节。 ⇒ 真实差异(会进日志/去重/hash),已改为对非法序列写回 RuneError,与旧行为逐值一致。 ## 诚实的收益结论:**基本没有** 改动后:mid_zh 189260(改前 186777)、long_zh 757746(改前 786416)、 mid_en 117343(改前 120422)。分配数 mid_en -40%、其余基本持平, **时间无实质改善**(部分场景还略慢)。 复查原因(不掩盖):重新做 CPU profile 后发现 **~25% 的样本是 runtime 锁/抢占**(unlock2 8.1% + lock2 6.8% + procyieldAsm 6.8% + asyncPreempt 5.4%),而 utf8/unicode 相关不足 20%。 且 GOMAXPROCS 敏感:1→375365ns、4→209922ns、12→189543ns ⇒ **大量时间花在调度与 GC 而非分词算术**。 ⇒ 结论:Go 侧微优化这条路**已到头**。真正的杠杆在别处: ① 提高 GOMAXPROCS/减少 GC 压力 ② 批量分词(降低每条输入的固定开销) ③ 减少送入模型的 token 量。三者都不是 C 能解决的。 改动本身保留(正确性等价、有 oracle 守护),但**不应据此宣称性能收益**。 与 C 化那几刀同一条纪律:没有数据支撑的优化不算优化。 验证:差分 oracle 6 组全过(含非法 UTF-8);providers/... 全绿。
This commit is contained in:
@ -18,6 +18,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"unicode/utf8"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
@ -120,7 +121,11 @@ func basicTokenize(text string) []string {
|
||||
cleaned := cleanText(text)
|
||||
var out []string
|
||||
for _, token := range strings.Fields(tokenizeChineseChars(cleaned)) {
|
||||
if len([]rune(token)) > maxInputCharsPerWord {
|
||||
// ★ 用 RuneCountInString 而不是 len([]rune(token)):
|
||||
// 后者为了**数一下长度**就把整个 token 转成 rune 切片 ⇒ 每个 token
|
||||
// 一次堆分配。而这个循环对每个词都跑,是分词器里最频繁的小动作。
|
||||
// 两者语义等价(都按 rune 计数,非法 UTF-8 每字节算一个 rune)。
|
||||
if utf8.RuneCountInString(token) > maxInputCharsPerWord {
|
||||
// 与 HF 一致:超长基本 token 直接丢弃(后续不会产出 UNK)。
|
||||
continue
|
||||
}
|
||||
@ -186,44 +191,92 @@ func stripAccents(text string) string {
|
||||
//
|
||||
// 注意 ASCII 段必须显式列出:'$' '+' '=' '^' '`' '|' '~' 属于 Sc/Sm/Sk,
|
||||
// 不是 Unicode P*,但它们也是标点(HF 用的是 ASCII 码点区间)。
|
||||
//
|
||||
// ★ 性能改动:把「累积 rune slice + 每次 string(cur)」换成
|
||||
// 一趟扫描,段以**字节区间**表示,最后一次 substring。
|
||||
// pprof 实测(mid_zh)本函数占 alloc_objects 的 33.5%。
|
||||
//
|
||||
// ★★ 但**非法 UTF-8 必须与旧实现逐值一致**:旧实现走 `[]rune(text)`,
|
||||
// 会把每个非法字节归一成 U+FFFD(`<60>`,3 字节);而纯字节切片会
|
||||
// **原样保留坏字节**。差分测试当场抓到这一分歧:
|
||||
// "\xbc\xef=..." → 旧 ["<22><>" ...] vs 新 ["\xbc\xef" ...]
|
||||
// 这是**真实缺陷**而非测量噪声:下游把 piece 当分词输入、也可能进日志,
|
||||
// 保留坏字节会让它进入本不该到达的地方(且 hash/去重会与旧行为不一致)。
|
||||
//
|
||||
// ⇒ 正确做法:**逐 rune 扫描**(utf8.DecodeRuneInString 对非法序列返回
|
||||
// (RuneError, 1),与 []rune 同语义),但**不预先把整串转成 rune slice**;
|
||||
// 对非 ASCII/非法字节的段,用 strings.Builder 写回 RuneError 的 UTF-8,
|
||||
// 从而与旧实现完全一致,同时省掉「整串 rune slice」那一块分配。
|
||||
func splitOnPunctuation(text string) []string {
|
||||
runes := []rune(text)
|
||||
var out []string
|
||||
var cur []rune
|
||||
var b strings.Builder
|
||||
b.Grow(len(text))
|
||||
hasBuf := false
|
||||
|
||||
flush := func() {
|
||||
if len(cur) > 0 {
|
||||
out = append(out, string(cur))
|
||||
cur = cur[:0]
|
||||
if hasBuf {
|
||||
out = append(out, b.String())
|
||||
b.Reset()
|
||||
hasBuf = false
|
||||
}
|
||||
}
|
||||
for _, r := range runes {
|
||||
|
||||
for i := 0; i < len(text); {
|
||||
r, size := utf8.DecodeRuneInString(text[i:])
|
||||
if isBERTPunctuation(r) {
|
||||
flush()
|
||||
out = append(out, string(r))
|
||||
i += size
|
||||
continue
|
||||
}
|
||||
cur = append(cur, r)
|
||||
// 普通字符:直接写原字节(与原实现 string([]rune) 等价)。
|
||||
// 非法序列:DecodeRuneInString 返回 RuneError,且 Go 的 []rune 也会
|
||||
// 产出 RuneError ⇒ 两者一致。
|
||||
if r == utf8.RuneError && size == 1 {
|
||||
b.WriteRune(utf8.RuneError)
|
||||
} else {
|
||||
b.WriteString(text[i : i+size])
|
||||
}
|
||||
hasBuf = true
|
||||
i += size
|
||||
}
|
||||
flush()
|
||||
return out
|
||||
}
|
||||
|
||||
// wordpiece 贪心最长匹配;整词任一段无法匹配则该词整体退化为 [UNK]。
|
||||
//
|
||||
// ★ 先定字节边界,再取一次 substring(而非每个候选都 string(runes[a:b])):
|
||||
// pprof 实测(mid_zh)本函数占 alloc_objects 的 29.5%,是第二大分配源。
|
||||
// 根因是内层循环**每轮候选都构造一个 string**:
|
||||
// piece := string(runes[start:end]) // "##"+piece 又是第二次分配
|
||||
// 而绝大多数候选都是未命中(要慢慢缩短 end),也就是**绝大多数
|
||||
// 分配都是浪费的**。
|
||||
// 改为:在原始字符串上按 rune 边界倒着推 end,只对**命中前最后一次**
|
||||
// 候选做一次 substring。于是每次匹配尝试从「2 次分配」降为 0 次,
|
||||
// 只有真正命中的那一段才分配。
|
||||
// 语义严格不变:仍然是最长前缀匹配、仍然对未命中整体退 [UNK]。
|
||||
func (t *Tokenizer) wordpiece(token string) []string {
|
||||
runes := []rune(token)
|
||||
if len(runes) > maxInputCharsPerWord {
|
||||
// 先建立 rune 边界表(单次分配,比每轮 substring 便宜得多)
|
||||
if utf8.RuneCountInString(token) > maxInputCharsPerWord {
|
||||
return []string{tokenUNK}
|
||||
}
|
||||
bounds := runeBounds(token)
|
||||
nr := len(bounds) - 1 // rune 个数
|
||||
var out []string
|
||||
start := 0
|
||||
for start < len(runes) {
|
||||
end := len(runes)
|
||||
var cur string
|
||||
start := 0 // rune 下标
|
||||
for start < nr {
|
||||
end := nr
|
||||
found := false
|
||||
var cur string
|
||||
for end > start {
|
||||
piece := string(runes[start:end])
|
||||
// 先查词表(用原串零拷贝切片构造 map key 仍需 string,
|
||||
// 但 Go 对 map[string] 的短 key 查找有优化,且这里
|
||||
// 只在**命中**时才真正保留;未命中的候选仍需构造 key)。
|
||||
word := token[bounds[start]:bounds[end]]
|
||||
piece := word
|
||||
if start > 0 {
|
||||
piece = "##" + piece
|
||||
piece = "##" + word
|
||||
}
|
||||
if _, ok := t.vocab[piece]; ok {
|
||||
cur = piece
|
||||
@ -241,6 +294,20 @@ func (t *Tokenizer) wordpiece(token string) []string {
|
||||
return out
|
||||
}
|
||||
|
||||
// runeBounds 返回 token 的 rune 边界字节偏移(长度 = rune 数 + 1)。
|
||||
//
|
||||
// 单次分配存边界,避免 wordpiece 内层循环反复切分字符串。
|
||||
func runeBounds(s string) []int {
|
||||
b := make([]int, 0, utf8.RuneCountInString(s)+1)
|
||||
for i := 0; i < len(s); {
|
||||
b = append(b, i)
|
||||
_, size := utf8.DecodeRuneInString(s[i:])
|
||||
i += size
|
||||
}
|
||||
b = append(b, len(s))
|
||||
return b
|
||||
}
|
||||
|
||||
func isASCII(s string) bool {
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] >= 0x80 {
|
||||
|
||||
177
providers/chineseclip/tokenizer_bench_test.go
Normal file
177
providers/chineseclip/tokenizer_bench_test.go
Normal file
@ -0,0 +1,177 @@
|
||||
package chineseclip
|
||||
|
||||
// tokenizer_bench_test.go —— 分词器热路径基准(判定 C 化是否值得)。
|
||||
//
|
||||
// ============================ 为什么不依赖真实模型 ============================
|
||||
// LoadTokenizer 只需要 vocab.txt + maxLength,不需要 ONNX 产物;
|
||||
// 本基准用**合成词表**(结构与真实词表同形:含 ## 续接前缀与 CJK 字符),
|
||||
// 于是无需 CHINESECLIP_MODEL_DIR 即可在任意机器复现。
|
||||
//
|
||||
// ★ 同时**不假设「C 更快」**,而是先测出成本分布:
|
||||
// 分词流水线有 5 段(cleanText / tokenizeChineseChars / splitOnPunctuation
|
||||
// / stripAccents+ToLower / wordpiece),哪一段占大头要**测**出来。
|
||||
// 三刀教训:C 化的收益判据必须先有数据支撑。
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/text/unicode/norm"
|
||||
)
|
||||
|
||||
// synthVocab 造一个与 BERT 中文词表同形的词表。
|
||||
// 结构还原要点:单字 + ## 续接 + 少量多字词 + 4 个特殊 token。
|
||||
func synthVocab() map[string]int32 {
|
||||
v := make(map[string]int32, 32768)
|
||||
add := func(p string) {
|
||||
if _, ok := v[p]; !ok {
|
||||
v[p] = int32(len(v))
|
||||
}
|
||||
}
|
||||
for _, s := range []string{tokenCLS, tokenSEP, tokenPAD, tokenUNK} {
|
||||
add(s)
|
||||
}
|
||||
// ASCII 词与 ## 续接
|
||||
words := []string{"hello", "world", "user", "query", "memory", "agent",
|
||||
"ing", "er", "ed", "s", "ly", "tion", "##ing", "##er", "##ed"}
|
||||
for _, w := range words {
|
||||
add(w)
|
||||
}
|
||||
// 常用单字(含中英)
|
||||
singles := []string{"的", "了", "是", "在", "我", "你", "他", "们", "这", "那",
|
||||
"a", "b", "c", "x", "y", "z", "0", "1", "2"}
|
||||
for _, s := range singles {
|
||||
add(s)
|
||||
}
|
||||
// 双字词(让 wordpiece 有机会一次命中)
|
||||
for i := 0; i < 512; i++ {
|
||||
add(string(rune('A'+i%26)) + string(rune('a'+i/26)))
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func benchTok(b *testing.B) *Tokenizer {
|
||||
b.Helper()
|
||||
return &Tokenizer{vocab: synthVocab(), maxLength: 512}
|
||||
}
|
||||
|
||||
// benchInputs 覆盖真实分布:短查询 / 中文长句 / 英文长文 / 混合 / 超长。
|
||||
var benchInputs = map[string]string{
|
||||
"short_zh": "用户询问了系统状态",
|
||||
"short_en": "what is the system status",
|
||||
"mid_zh": strings.Repeat("这是一段中文文本,用于测试分词器的吞吐。", 10),
|
||||
"mid_en": strings.Repeat("the quick brown fox jumps over the lazy dog. ", 10),
|
||||
"mixed": strings.Repeat("记忆 memory 检索 recall 上下文 context 注入 inject。", 8),
|
||||
"long_zh": strings.Repeat("长文本。", 200),
|
||||
"punct_heavy": strings.Repeat("你好,世界!这是一个测试。", 20),
|
||||
}
|
||||
|
||||
// BenchmarkTokenizerEncode 整体分词(Encode 全流程)。
|
||||
func BenchmarkTokenizerEncode(b *testing.B) {
|
||||
tok := benchTok(b)
|
||||
for name, in := range benchInputs {
|
||||
b.Run(name, func(b *testing.B) {
|
||||
b.SetBytes(int64(len(in)))
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
tok.Encode(in)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTokenizerStages 分解流水线各段,用于定位真正的热点。
|
||||
//
|
||||
// ★ 判据:若 wordpiece 占比远高于其它段,则「O(n²) 字符串分配」是主因;
|
||||
// 若 basicTokenize 的分配占比高,则 cleanText/tokenizeChineseChars 的
|
||||
// strings.Builder 往返是主因。两者处方完全不同,不能凭直觉断言。
|
||||
func BenchmarkTokenizerStages(b *testing.B) {
|
||||
tok := benchTok(b)
|
||||
for _, name := range []string{"mid_zh", "mid_en"} {
|
||||
in := benchInputs[name]
|
||||
b.Run(name+"/basicTokenize", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
basicTokenize(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/cleanText", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
cleanText(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/tokenizeChineseChars", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
tokenizeChineseChars(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/stripAccents", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
stripAccents(strings.ToLower(in))
|
||||
}
|
||||
})
|
||||
b.Run(name+"/wordpiece_all", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for _, tok2 := range basicTokenize(in) {
|
||||
tok.wordpiece(tok2)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkWordpieceSingle 单独压 wordpiece(已知 O(n²) 分配的那个)。
|
||||
func BenchmarkWordpieceSingle(b *testing.B) {
|
||||
tok := benchTok(b)
|
||||
// 长 token 触发更多次回退(end 递减)
|
||||
cases := map[string]string{
|
||||
"cjk_1": "学",
|
||||
"cjk_2": "学习",
|
||||
"cjk_4": "学习机器学习",
|
||||
"ascii_8": "unbelievable",
|
||||
"mixed_6": "机器learning",
|
||||
}
|
||||
for name, w := range cases {
|
||||
b.Run(name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
tok.wordpiece(w)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTokenizerSplitPuncAndToLower 补测两段之前没单独量的成本。
|
||||
func BenchmarkTokenizerSplitPuncAndToLower(b *testing.B) {
|
||||
for _, name := range []string{"mid_zh", "mid_en", "punct_heavy"} {
|
||||
in := benchInputs[name]
|
||||
b.Run(name+"/splitOnPunctuation", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
splitOnPunctuation(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/ToLower", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = strings.ToLower(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/Fields", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = strings.Fields(in)
|
||||
}
|
||||
})
|
||||
b.Run(name+"/NFD", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = norm.NFD.String(in)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
290
providers/chineseclip/tokenizer_diff_test.go
Normal file
290
providers/chineseclip/tokenizer_diff_test.go
Normal file
@ -0,0 +1,290 @@
|
||||
package chineseclip
|
||||
|
||||
// tokenizer_diff_test.go —— 新实现 vs **原实现(内联为参照 oracle)** 的差分等价测试。
|
||||
//
|
||||
// ============================ 为什么必须有这个文件 ============================
|
||||
// 本轮我把 splitOnPunctuation 与 wordpiece 从「[]rune + 每轮 substring」
|
||||
// 改成「字节边界 + 一次 substring」。目标是纯性能,语义必须**逐值不变**。
|
||||
//
|
||||
// 而本机**没有真实模型产物**(CHINESECLIP_MODEL_DIR 未设),
|
||||
// 唯一的权威对照 TestTokenizerMatchesOfficialReference会 **SKIP** ——
|
||||
// 也就是说:仅靠现有测试,我的重写是**没有被有效验证**的。
|
||||
//
|
||||
// 故这里把**原实现**原样内联为 oracle,用同一批输入逐值比对。
|
||||
// 这与本仓 C 化那几刀同一条纪律:
|
||||
// 「没有对照的优化,只是感觉而不是证据」。
|
||||
//
|
||||
// oracle 是**冻结的旧代码**(照抄改动前的实现),不得随主实现演进而修改——
|
||||
// 否则它就失去了「参照」的意义。
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// oracle:改动前的 splitOnPunctuation / wordpiece(原样冻结)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func oracleSplitOnPunctuation(text string) []string {
|
||||
runes := []rune(text)
|
||||
var out []string
|
||||
var cur []rune
|
||||
flush := func() {
|
||||
if len(cur) > 0 {
|
||||
out = append(out, string(cur))
|
||||
cur = cur[:0]
|
||||
}
|
||||
}
|
||||
for _, r := range runes {
|
||||
if isBERTPunctuation(r) {
|
||||
flush()
|
||||
out = append(out, string(r))
|
||||
continue
|
||||
}
|
||||
cur = append(cur, r)
|
||||
}
|
||||
flush()
|
||||
return out
|
||||
}
|
||||
|
||||
func oracleWordpiece(vocab map[string]int32, token string) []string {
|
||||
runes := []rune(token)
|
||||
if len(runes) > maxInputCharsPerWord {
|
||||
return []string{tokenUNK}
|
||||
}
|
||||
var out []string
|
||||
start := 0
|
||||
for start < len(runes) {
|
||||
end := len(runes)
|
||||
var cur string
|
||||
found := false
|
||||
for end > start {
|
||||
piece := string(runes[start:end])
|
||||
if start > 0 {
|
||||
piece = "##" + piece
|
||||
}
|
||||
if _, ok := vocab[piece]; ok {
|
||||
cur = piece
|
||||
found = true
|
||||
break
|
||||
}
|
||||
end--
|
||||
}
|
||||
if !found {
|
||||
return []string{tokenUNK}
|
||||
}
|
||||
out = append(out, cur)
|
||||
start = end
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func oracleBasicTokenize(text string) []string {
|
||||
cleaned := oracleCleanText(text)
|
||||
var out []string
|
||||
for _, token := range strings.Fields(oracleTokenizeChineseChars(cleaned)) {
|
||||
if len([]rune(token)) > maxInputCharsPerWord {
|
||||
continue
|
||||
}
|
||||
stripped := stripAccents(strings.ToLower(token))
|
||||
out = append(out, oracleSplitOnPunctuation(stripped)...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// oracleCleanText / oracleTokenizeChineseChars 直接复用主实现中**未被改动**的
|
||||
// 函数(它们本轮没动,故无需再抄一份,抄了反而会有漂移风险)。
|
||||
func oracleCleanText(text string) string { return cleanText(text) }
|
||||
|
||||
func oracleTokenizeChineseChars(text string) string { return tokenizeChineseChars(text) }
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 差分测试
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// diffInputs 覆盖各语言/标点/空白/emoji/组合字符/超长词。
|
||||
var diffInputs = []string{
|
||||
"",
|
||||
"a",
|
||||
"你好",
|
||||
"你好,世界!",
|
||||
"hello world",
|
||||
"hello, world!",
|
||||
"用户询问了系统状态",
|
||||
"这是一段中文文本,用于测试分词器的吞吐。",
|
||||
"the quick brown fox jumps over the lazy dog.",
|
||||
"记忆 memory 检索 recall 上下文 context 注入 inject。",
|
||||
"混合Mixed中英English文本text。",
|
||||
"標點測試:;、()《》「」",
|
||||
"emoji 😀 与中文混合",
|
||||
"combin\u0301ing", // 组合音标
|
||||
"café naïve résumé", // 预组合
|
||||
"a" + strings.Repeat("b", 300), // 超长 ASCII 词(> maxInputCharsPerWord)
|
||||
"字" + strings.Repeat("长", 300),
|
||||
" 多个 空格\t制表\n换行 ",
|
||||
"$+=^`|~ 符号",
|
||||
"#全角#ABC", // 全角
|
||||
"1234567890",
|
||||
"UPPER lower MiXeD",
|
||||
"无标点长句onetwothreefour",
|
||||
"\u0000\u0001控制符",
|
||||
"a,,b。c!d?e;f:g",
|
||||
strings.Repeat("词。", 100),
|
||||
}
|
||||
|
||||
func TestTokenizerDiff_SplitOnPunctuation(t *testing.T) {
|
||||
for _, in := range diffInputs {
|
||||
got := splitOnPunctuation(in)
|
||||
want := oracleSplitOnPunctuation(in)
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("splitOnPunctuation 分歧 %q:\n got %q\n want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenizerDiff_Wordpiece(t *testing.T) {
|
||||
vocab := synthVocab()
|
||||
tok := &Tokenizer{vocab: vocab, maxLength: 512}
|
||||
// 单独的词(含超长、含 CJK、含 ASCII、含未登录词)
|
||||
words := []string{
|
||||
"hello", "world", "helloing", "unbelievable", "abc", "a",
|
||||
"学", "学习", "学习机器学习", "机器learning", "未知词汇",
|
||||
strings.Repeat("x", 300), strings.Repeat("学", 300),
|
||||
"café", "caféing", "aaaaaaa",
|
||||
"", "##x", "1234567890",
|
||||
}
|
||||
for _, w := range words {
|
||||
got := tok.wordpiece(w)
|
||||
want := oracleWordpiece(vocab, w)
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("wordpiece 分歧 %q:\n got %v\n want %v", w, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenizerDiff_BasicTokenize(t *testing.T) {
|
||||
for _, in := range diffInputs {
|
||||
got := basicTokenize(in)
|
||||
want := oracleBasicTokenize(in)
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("basicTokenize 分歧 %q:\n got %v\n want %v", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenizerDiff_Encode 端到端:整条 Encode(含特殊 token 与 padding)。
|
||||
func TestTokenizerDiff_Encode(t *testing.T) {
|
||||
vocab := synthVocab()
|
||||
tok := &Tokenizer{vocab: vocab, maxLength: 512}
|
||||
for _, in := range diffInputs {
|
||||
got, maskGot := tok.Encode(in)
|
||||
|
||||
// oracle 路径:用冻结的 basicTokenize + wordpiece 重建 Encode
|
||||
pieces := []string{}
|
||||
for _, basic := range oracleBasicTokenize(in) {
|
||||
pieces = append(pieces, oracleWordpiece(vocab, basic)...)
|
||||
}
|
||||
if limit := 512 - 2; len(pieces) > limit {
|
||||
pieces = pieces[:limit]
|
||||
}
|
||||
want := []int64{int64(vocab[tokenCLS])}
|
||||
for _, p := range pieces {
|
||||
want = append(want, int64(vocab[p]))
|
||||
}
|
||||
want = append(want, int64(vocab[tokenSEP]))
|
||||
for len(want) < 512 {
|
||||
want = append(want, int64(vocab[tokenPAD]))
|
||||
}
|
||||
wantMask := make([]int64, 0, 512)
|
||||
n := len(pieces) + 2
|
||||
for i := 0; i < n; i++ {
|
||||
wantMask = append(wantMask, 1)
|
||||
}
|
||||
for len(wantMask) < 512 {
|
||||
wantMask = append(wantMask, 0)
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("Encode ids 分歧 %q", in)
|
||||
}
|
||||
if !reflect.DeepEqual(maskGot, wantMask) {
|
||||
t.Errorf("Encode mask 分歧 %q", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenizerDiff_RandomBytes 随机字节:非法 UTF-8 是分词器最容易分叉的输入。
|
||||
//
|
||||
// ★ 为什么必须测这个:字节层面扫描(新实现)与 rune 层面扫描(旧实现)在
|
||||
// **非法 UTF-8** 上的行为最容易不同 ——
|
||||
// · []rune(s) 把非法字节变成 U+FFFD(每个坏字节一个)
|
||||
// · utf8.DecodeRuneInString 返回 (RuneError, 1) 并前进 1 字节
|
||||
// 两者语义应当一致,但「应当」不是证据。
|
||||
func TestTokenizerDiff_RandomBytes(t *testing.T) {
|
||||
rng := newSeededRand(20260926)
|
||||
alphabet := []byte("ab ,.!?中文。,!?$+=^`|~\xff\xfe\x80\xc3\xe4\t\n")
|
||||
for iter := 0; iter < 20000; iter++ {
|
||||
n := rng.Intn(40)
|
||||
buf := make([]byte, n)
|
||||
for i := range buf {
|
||||
buf[i] = alphabet[rng.Intn(len(alphabet))]
|
||||
}
|
||||
in := string(buf)
|
||||
|
||||
if got, want := splitOnPunctuation(in), oracleSplitOnPunctuation(in); !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("随机字节 splitOnPunctuation 分歧 %q:\n got %q\n want %q", in, got, want)
|
||||
}
|
||||
if got, want := basicTokenize(in), oracleBasicTokenize(in); !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("随机字节 basicTokenize 分歧 %q:\n got %v\n want %v", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenizerDiff_FuzzishRunes 随机 rune(合法 UTF-8 但内容任意)。
|
||||
func TestTokenizerDiff_FuzzishRunes(t *testing.T) {
|
||||
rng := newSeededRand(777)
|
||||
var sb strings.Builder
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
sb.Reset()
|
||||
n := rng.Intn(30)
|
||||
for i := 0; i < n; i++ {
|
||||
sb.WriteRune(rune(rng.Intn(0x2000)))
|
||||
}
|
||||
in := sb.String()
|
||||
if got, want := splitOnPunctuation(in), oracleSplitOnPunctuation(in); !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("随机 rune 分歧 %q:\n got %q\n want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 辅助
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
type seededRand struct{ s uint64 }
|
||||
|
||||
func newSeededRand(seed uint64) *seededRand { return &seededRand{s: seed | 1} }
|
||||
|
||||
func (r *seededRand) next() uint64 {
|
||||
r.s ^= r.s << 13
|
||||
r.s ^= r.s >> 7
|
||||
r.s ^= r.s << 17
|
||||
return r.s
|
||||
}
|
||||
|
||||
func (r *seededRand) Intn(n int) int {
|
||||
if n <= 0 {
|
||||
return 0
|
||||
}
|
||||
return int(r.next() % uint64(n))
|
||||
}
|
||||
|
||||
// 确保 oracle 与主实现对 isBertPunctuation 的使用一致(防有人改了判定)。
|
||||
var (
|
||||
_ = unicode.IsPunct
|
||||
_ = utf8.RuneError
|
||||
)
|
||||
Reference in New Issue
Block a user