Skip to content

Commit 8cab673

Browse files
committed
feat(kb): bge-small-zh BERT WordPiece 分词器(REQ-259 地基二)
1 parent 244ea92 commit 8cab673

2 files changed

Lines changed: 196 additions & 0 deletions

File tree

Lines changed: 194 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,194 @@
1+
//! bge-small-zh-v1.5 的 BERT WordPiece 分词(REQ-259,v0.19.5)。
2+
//!
3+
//! @ai-context: 内嵌 ONNX 引擎需要与训练一致的输入(input_ids/attention_mask)。
4+
//! 官方句向量模型用 WordPiece + vocab.txt——中文语料为字符级词表
5+
//! 全覆盖;本模块实现**最小 WordPiece**(不含 lowercase/重音剥离:
6+
//! zh 模型词表为大写敏感无关字符级;ASCII 按字面匹配、未命中回退
7+
//! [UNK]——诚实降级而非猜测)。
8+
//! @ai-context: 契约:`[CLS] tokens... [SEP]`,max_len=512(超长截断——bge
9+
//! 训练窗 512;kb 切块 ≤800 字符硬切已留安全余量),pad 到
10+
//! max_len(batch 同长)。词表从模型目录 vocab.txt 加载(与模型
11+
//! 同分发);加载失败 → 引擎不可用(外层诚实报错,检索自动降级)。
12+
//! @ai-context: 特殊标记取自词表查找(防御换名);词表缺 [UNK]/[CLS]/[SEP]
13+
//! → 加载失败(模型不完整,宁缺勿错)。
14+
//! 注意:引擎模块(kb_embed_onnx)接线前本模块无生产调用——dead_code 临时
15+
//! 豁免,引擎接线轮必须移除本属性(TODO REQ-259)。
16+
#![allow(dead_code)]
17+
18+
use std::collections::HashMap;
19+
use std::path::Path;
20+
21+
/// bge-small-zh-v1.5 上下文窗(模型 max_position_embeddings)
22+
pub const MAX_LEN: usize = 512;
23+
24+
/// 默认特殊标记名(词表内实际名字为准——查找失败即加载失败)
25+
const PAD: &str = "[PAD]";
26+
const CLS: &str = "[CLS]";
27+
const SEP: &str = "[SEP]";
28+
const UNK: &str = "[UNK]";
29+
30+
/// 分词器(vocab.txt 加载后不可变——Send+Sync,可跨任务线程复用)
31+
#[derive(Debug, Clone)]
32+
pub struct BertTokenizer {
33+
word_to_id: HashMap<String, u32>,
34+
/// 词表条目文本(id → token;构建 ## 续接查找)
35+
tokens: Vec<String>,
36+
pad_id: u32,
37+
cls_id: u32,
38+
sep_id: u32,
39+
unk_id: u32,
40+
}
41+
42+
/// 一次编码结果(ndarray 前的一维数组——由引擎转 batch 矩阵)
43+
#[derive(Debug, Clone, PartialEq)]
44+
pub struct Encoded {
45+
pub input_ids: Vec<i64>,
46+
pub attention_mask: Vec<i64>,
47+
}
48+
49+
impl BertTokenizer {
50+
/// 从 vocab.txt 加载(每行一个 token;带 `##` 续接与特殊标记校验)。
51+
pub fn load(path: &Path) -> Result<Self, String> {
52+
let raw = std::fs::read_to_string(path).map_err(|e| format!("词表读取失败: {e}"))?;
53+
let tokens: Vec<String> = raw.lines().map(str::trim).filter(|l| !l.is_empty()).map(str::to_string).collect();
54+
if tokens.is_empty() {
55+
return Err("词表为空".to_string());
56+
}
57+
let find = |name: &str| -> Option<u32> {
58+
tokens.iter().position(|t| t == name).map(|i| i as u32)
59+
};
60+
let pad_id = find(PAD).ok_or_else(|| format!("词表缺少 {PAD}"))?;
61+
let cls_id = find(CLS).ok_or_else(|| format!("词表缺少 {CLS}"))?;
62+
let sep_id = find(SEP).ok_or_else(|| format!("词表缺少 {SEP}"))?;
63+
let unk_id = find(UNK).ok_or_else(|| format!("词表缺少 {UNK}"))?;
64+
let word_to_id: HashMap<String, u32> =
65+
tokens.iter().enumerate().map(|(i, t)| (t.clone(), i as u32)).collect();
66+
Ok(Self { word_to_id, tokens, pad_id, cls_id, sep_id, unk_id })
67+
}
68+
69+
/// 编码单条文本 → [CLS] … [SEP](截断到 MAX_LEN-2 个内容 token,pad 满窗)。
70+
pub fn encode(&self, text: &str) -> Encoded {
71+
let mut ids: Vec<i64> = vec![self.cls_id as i64];
72+
let content_cap = MAX_LEN.saturating_sub(2);
73+
for piece in self.word_pieces(text).take(content_cap) {
74+
ids.push(self.word_to_id.get(&piece).copied().unwrap_or(self.unk_id) as i64);
75+
}
76+
ids.push(self.sep_id as i64);
77+
let len = ids.len();
78+
ids.resize(MAX_LEN, self.pad_id as i64);
79+
let mut mask = vec![0i64; MAX_LEN];
80+
mask[..len].fill(1);
81+
Encoded { input_ids: ids, attention_mask: mask }
82+
}
83+
84+
/// 最小 WordPiece:按 Unicode 标量切分,逐字符最长匹配(≤100 步内)
85+
/// + `##` 续接词表条目。
86+
fn word_pieces<'a>(&'a self, text: &'a str) -> impl Iterator<Item = String> + 'a {
87+
let chars: Vec<char> = text.chars().collect();
88+
let mut out: Vec<String> = Vec::new();
89+
let mut i = 0usize;
90+
while i < chars.len() {
91+
// 词表条目可能跨多字符(常见中文词/符号)——最长匹配(≤100 步)
92+
let mut matched: Option<(usize, String)> = None; // (消费字符数, token)
93+
let max_step = (chars.len().saturating_sub(i)).min(100);
94+
for step in (1..=max_step).rev() {
95+
let cand: String = chars[i..i + step].iter().collect();
96+
if self.word_to_id.contains_key(&cand) {
97+
matched = Some((step, cand));
98+
break;
99+
}
100+
}
101+
let (consumed, token) = match matched {
102+
Some(m) => m,
103+
None => {
104+
// 单字符未命中 → 尝试 ## 续接形式(模型词表可能只存 ##xx)
105+
let ch: String = chars[i..i + 1].iter().collect();
106+
let cont = format!("##{ch}");
107+
if self.word_to_id.contains_key(&cont) {
108+
out.push(cont);
109+
} else {
110+
out.push(self.unk_id_token_name());
111+
}
112+
i += 1;
113+
continue;
114+
}
115+
};
116+
out.push(token);
117+
i += consumed;
118+
}
119+
out.into_iter()
120+
}
121+
122+
/// UNK 名称(仅测试/诊断用——真实 token 由 unk_id 索引保证一致)
123+
fn unk_id_token_name(&self) -> String {
124+
self.tokens.get(self.unk_id as usize).cloned().unwrap_or_else(|| UNK.to_string())
125+
}
126+
}
127+
128+
#[cfg(test)]
129+
mod tests {
130+
use super::*;
131+
132+
/// 构造迷你词表(常见字 + [CLS]/[SEP]/[UNK]/[PAD] + 一个 ## 续接)
133+
fn write_vocab(path: &std::path::Path) {
134+
let vocab = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "学", "习", "编", "程", "##的", "今天", "好"];
135+
std::fs::write(path, vocab.join("\n") + "\n").unwrap();
136+
}
137+
138+
fn temp_vocab(tag: &str) -> BertTokenizer {
139+
// 用例级唯一目录(并行测试共享 pid 目录会互删互踩——race 源)
140+
static SEQ: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
141+
let n = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
142+
let dir = std::env::temp_dir().join(format!("entropy-bpe-{tag}-{}-{n}", std::process::id()));
143+
let _ = std::fs::remove_dir_all(&dir);
144+
std::fs::create_dir_all(&dir).unwrap();
145+
let path = dir.join("vocab.txt");
146+
write_vocab(&path);
147+
BertTokenizer::load(&path).unwrap()
148+
}
149+
150+
#[test]
151+
fn load_rejects_incomplete_vocab() {
152+
let dir = std::env::temp_dir().join(format!("entropy-bpe-bad-{}", std::process::id()));
153+
std::fs::create_dir_all(&dir).unwrap();
154+
let p = dir.join("vocab.txt");
155+
std::fs::write(&p, "学\n习\n").unwrap();
156+
assert!(BertTokenizer::load(&p).is_err(), "缺特殊标记必须加载失败");
157+
assert!(BertTokenizer::load(&dir.join("none.txt")).is_err());
158+
}
159+
160+
#[test]
161+
fn encode_has_cls_sep_pad_and_mask() {
162+
let tok = temp_vocab("cls");
163+
let e = tok.encode("今天学习编程");
164+
assert_eq!(e.input_ids.len(), MAX_LEN);
165+
assert_eq!(e.input_ids[0], tok.cls_id as i64);
166+
// 尾随 SEP(紧跟在内容后)
167+
let first_pad = e.input_ids.iter().position(|&v| v == tok.pad_id as i64).unwrap_or(MAX_LEN);
168+
assert_eq!(e.input_ids[first_pad - 1], tok.sep_id as i64, "SEP 在内容后、pad 前");
169+
// mask:1..=内容+CLS+SEP 段,其余 0
170+
let ones = e.attention_mask.iter().filter(|&&m| m == 1).count();
171+
assert_eq!(ones, first_pad, "mask 段与真实 token 数一致(截断窗内)");
172+
// 词表内字应编码为自身 id 而非 UNK
173+
assert!(e.input_ids[1..first_pad - 1].iter().all(|&v| v != tok.unk_id as i64), "词表内字不得回落 UNK");
174+
}
175+
176+
#[test]
177+
fn unknown_char_falls_back_to_unk_without_panic() {
178+
let tok = temp_vocab("unk");
179+
let e = tok.encode("学习 🚀 编程");
180+
let unk = tok.unk_id as i64;
181+
assert!(e.input_ids.contains(&unk), "未登录字符 → UNK 诚实降级");
182+
assert_eq!(e.input_ids.len(), MAX_LEN);
183+
}
184+
185+
#[test]
186+
fn long_text_truncated_to_window() {
187+
let tok = temp_vocab("long");
188+
let long = "学".repeat(3000);
189+
let e = tok.encode(&long);
190+
// 内容 token 数 = MAX_LEN-2(CLS+SEP 占位)
191+
let content_len = e.attention_mask.iter().filter(|&&m| m == 1).count() - 2;
192+
assert_eq!(content_len, MAX_LEN - 2);
193+
}
194+
}

‎app/src-tauri/src/lib.rs‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -264,6 +264,8 @@ mod commands_ai_chat_kb;
264264
mod kb_discovery;
265265
// REQ-259(v0.19.5):kb 语义索引 embedding 契约与纯函数(编解码/cosine top-K)
266266
mod kb_embed;
267+
// REQ-259(v0.19.5):bge-small-zh BERT WordPiece 分词(vocab.txt 加载/编码)
268+
mod kb_embed_tokenizer;
267269
mod commands_kb_discovery;
268270
mod concept_weakness;
269271
mod goal_interview;

0 commit comments

Comments
 (0)