|
| 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 | +} |
0 commit comments