diff --git a/providers/chineseclip/tokenizer.go b/providers/chineseclip/tokenizer.go index 4894a2e..a59140c 100644 --- a/providers/chineseclip/tokenizer.go +++ b/providers/chineseclip/tokenizer.go @@ -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(`�`,3 字节);而纯字节切片会 +// **原样保留坏字节**。差分测试当场抓到这一分歧: +// "\xbc\xef=..." → 旧 ["��" ...] 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 { diff --git a/providers/chineseclip/tokenizer_bench_test.go b/providers/chineseclip/tokenizer_bench_test.go new file mode 100644 index 0000000..00c7af5 --- /dev/null +++ b/providers/chineseclip/tokenizer_bench_test.go @@ -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) + } + }) + } +} diff --git a/providers/chineseclip/tokenizer_diff_test.go b/providers/chineseclip/tokenizer_diff_test.go new file mode 100644 index 0000000..bb4e554 --- /dev/null +++ b/providers/chineseclip/tokenizer_diff_test.go @@ -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 +)