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:
JianFeeeee
2026-09-26 11:44:49 +08:00
parent 4ead049a52
commit 15e5e87fbf
3 changed files with 550 additions and 16 deletions

View 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
)