|
| 1 | +//! bge-small-zh-v1.5 的 ONNX embedding 引擎(REQ-259,v0.19.5)。 |
| 2 | +//! |
| 3 | +//! @ai-context: ort 2.0.0-rc.13 会话封装(CPU EP;与 oar-ocr 共用同一份 |
| 4 | +//! onnxruntime.dll——.cargo/config.toml ORT_LIB_LOCATION 注入)。 |
| 5 | +//! 模型目录约定:`<data_dir>/models/embedding/bge-small-zh-v1.5/` |
| 6 | +//! 内含 `model_quantized.onnx`(或 model.onnx)+ `vocab.txt`。 |
| 7 | +//! @ai-context: 编码=WordPiece(kb_embed_tokenizer)→ 批推理 → 取 |
| 8 | +//! last_hidden_state 的 [CLS] 行 → L2 归一。任何加载/推理失败 → |
| 9 | +//! Err(外层按"未嵌"降级 FTS-only,状态命令如实暴露原因)。 |
| 10 | +//! @ai-context: rc13 API 注意:Session::run 需 &mut self(内部 Mutex 串行化); |
| 11 | +//! 输入张量经 ort::value::Tensor::from_array 构造;输出读取用 |
| 12 | +//! Value::try_extract_tensor::<f32>() → (&Shape, &[f32])。 |
| 13 | +
|
| 14 | +use std::path::Path; |
| 15 | +use std::sync::{Arc, Mutex}; |
| 16 | + |
| 17 | +use ort::session::Session; |
| 18 | +use ort::value::Tensor; |
| 19 | + |
| 20 | +use crate::kb_embed::EmbeddingEngine; |
| 21 | +use crate::kb_embed_tokenizer::{BertTokenizer, Encoded, MAX_LEN}; |
| 22 | + |
| 23 | +/// bge-small-zh-v1.5 输出维度(模型卡固定 512;运行期以数据长度复核) |
| 24 | +pub const BGE_DIM: usize = 512; |
| 25 | +/// 单次推理最大批(内存有界——全量重嵌分批走) |
| 26 | +pub const MAX_BATCH: usize = 32; |
| 27 | + |
| 28 | +/// 模型文件候选名(量化优先;社区导出惯例两态) |
| 29 | +const MODEL_CANDIDATES: [&str; 4] = [ |
| 30 | + "model_quantized.onnx", |
| 31 | + "model.onnx", |
| 32 | + "onnx/model_quantized.onnx", |
| 33 | + "onnx/model.onnx", |
| 34 | +]; |
| 35 | + |
| 36 | +/// ort 会话(run 需 &mut —— Mutex 串行化;Send 由 Session 保证) |
| 37 | +struct OnnxSession { |
| 38 | + inner: Mutex<Session>, |
| 39 | +} |
| 40 | + |
| 41 | +impl OnnxSession { |
| 42 | + fn open(path: &Path) -> Result<Self, String> { |
| 43 | + let session = Session::builder() |
| 44 | + .map_err(|e| format!("ort 初始化失败: {e}"))? |
| 45 | + .commit_from_file(path) |
| 46 | + .map_err(|e| format!("模型加载失败({path:?}): {e}"))?; |
| 47 | + Ok(Self { inner: Mutex::new(session) }) |
| 48 | + } |
| 49 | + |
| 50 | + /// 批推理 → 每行 [CLS] 行向量(len=BGE_DIM;形状不符/取锁失败 → Err) |
| 51 | + fn embed_batch(&self, encs: &[Encoded]) -> Result<Vec<Vec<f32>>, String> { |
| 52 | + let n = encs.len(); |
| 53 | + if n == 0 { |
| 54 | + return Ok(Vec::new()); |
| 55 | + } |
| 56 | + let flat_ids: Vec<i64> = encs.iter().flat_map(|e| e.input_ids.iter().copied()).collect(); |
| 57 | + let flat_mask: Vec<i64> = encs.iter().flat_map(|e| e.attention_mask.iter().copied()).collect(); |
| 58 | + let ids = Tensor::<i64>::from_array(([n, MAX_LEN], flat_ids)) |
| 59 | + .map_err(|e| format!("输入张量构造失败: {e}"))?; |
| 60 | + let mask = Tensor::<i64>::from_array(([n, MAX_LEN], flat_mask)) |
| 61 | + .map_err(|e| format!("输入张量构造失败: {e}"))?; |
| 62 | + let mut sess = self.inner.lock().map_err(|_| "ort 会话锁中毒".to_string())?; |
| 63 | + let outputs = sess |
| 64 | + .run(ort::inputs!["input_ids" => ids, "attention_mask" => mask]) |
| 65 | + .map_err(|e| format!("推理失败: {e}"))?; |
| 66 | + // last_hidden_state 约定为首输出:形状 (n, seq, hidden)——flat 切片直读 |
| 67 | + let (_shape, data) = outputs[0] |
| 68 | + .try_extract_tensor::<f32>() |
| 69 | + .map_err(|e| format!("输出取张量失败: {e}"))?; |
| 70 | + let per_row = data.len() / n; |
| 71 | + if per_row == 0 || !per_row.is_multiple_of(MAX_LEN) { |
| 72 | + return Err(format!("输出长度意外: {}(预期 n×seq×hidden)", data.len())); |
| 73 | + } |
| 74 | + let hidden = per_row / MAX_LEN; |
| 75 | + if hidden != BGE_DIM { |
| 76 | + return Err(format!("模型 hidden 维度 {hidden} ≠ 预期 {BGE_DIM}")); |
| 77 | + } |
| 78 | + let mut out = Vec::with_capacity(n); |
| 79 | + for i in 0..n { |
| 80 | + let base = i * per_row; |
| 81 | + out.push(data[base..base + hidden].to_vec()); // [CLS] 行 = 序列首 token |
| 82 | + } |
| 83 | + Ok(out) |
| 84 | + } |
| 85 | +} |
| 86 | + |
| 87 | +/// ONNX embedding 引擎(加载即校验模型文件与词表;推理形状首用即校验)。 |
| 88 | +pub struct OnnxEmbedding { |
| 89 | + session: Arc<OnnxSession>, |
| 90 | + tokenizer: BertTokenizer, |
| 91 | +} |
| 92 | + |
| 93 | +impl OnnxEmbedding { |
| 94 | + /// 从模型目录加载(模型与 vocab 同目录;缺任一 → Err 携带可诊断原因)。 |
| 95 | + pub fn try_load(dir: &Path) -> Result<Self, String> { |
| 96 | + let model = MODEL_CANDIDATES |
| 97 | + .iter() |
| 98 | + .map(|name| dir.join(name)) |
| 99 | + .find(|p| p.is_file()) |
| 100 | + .ok_or_else(|| format!("模型文件缺失:{dir:?}(需 model_quantized.onnx 或 model.onnx)"))?; |
| 101 | + let tokenizer = BertTokenizer::load(&dir.join("vocab.txt"))?; |
| 102 | + let session = Arc::new(OnnxSession::open(&model)?); |
| 103 | + Ok(Self { session, tokenizer }) |
| 104 | + } |
| 105 | +} |
| 106 | + |
| 107 | +/// L2 归一(零向量保持原样——余弦侧已有零守卫) |
| 108 | +fn l2_normalize(mut v: Vec<f32>) -> Vec<f32> { |
| 109 | + let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt(); |
| 110 | + if norm > 0.0 { |
| 111 | + for x in v.iter_mut() { |
| 112 | + *x /= norm; |
| 113 | + } |
| 114 | + } |
| 115 | + v |
| 116 | +} |
| 117 | + |
| 118 | +impl EmbeddingEngine for OnnxEmbedding { |
| 119 | + fn dims(&self) -> Option<usize> { |
| 120 | + Some(BGE_DIM) |
| 121 | + } |
| 122 | + fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, String> { |
| 123 | + let mut out = Vec::with_capacity(texts.len()); |
| 124 | + for batch in texts.chunks(MAX_BATCH) { |
| 125 | + let encs: Vec<Encoded> = batch.iter().map(|t| self.tokenizer.encode(t)).collect(); |
| 126 | + for row in self.session.embed_batch(&encs)? { |
| 127 | + out.push(l2_normalize(row)); |
| 128 | + } |
| 129 | + } |
| 130 | + Ok(out) |
| 131 | + } |
| 132 | +} |
| 133 | + |
| 134 | +#[cfg(test)] |
| 135 | +mod tests { |
| 136 | + use super::*; |
| 137 | + |
| 138 | + #[test] |
| 139 | + fn l2_normalize_zero_and_unit_vectors() { |
| 140 | + let v = l2_normalize(vec![3.0, 4.0]); |
| 141 | + assert!((v[0] - 0.6).abs() < 1e-5 && (v[1] - 0.8).abs() < 1e-5); |
| 142 | + assert_eq!(l2_normalize(vec![0.0, 0.0]), vec![0.0, 0.0]); |
| 143 | + } |
| 144 | + |
| 145 | + #[test] |
| 146 | + fn try_load_missing_model_dir_is_honest_error() { |
| 147 | + let missing = std::env::temp_dir().join(format!("entropy-onnx-missing-{}", std::process::id())); |
| 148 | + let err = match OnnxEmbedding::try_load(&missing) { |
| 149 | + Ok(_) => panic!("缺模型必须失败"), |
| 150 | + Err(e) => e, |
| 151 | + }; |
| 152 | + assert!(err.contains("模型文件缺失"), "无模型必须如实报错: {err}"); |
| 153 | + } |
| 154 | +} |
0 commit comments