Files
HomeAgent/providers/chineseclip/tokenizer_diff_test.go
JianFeeeee 84d2d7c313 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/... 全绿。
2026-09-26 11:44:49 +08:00

291 lines
8.8 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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
)