Skip to content

Commit c57ce0c

Browse files
committed
feat(kb): OnnxEmbedding 引擎与状态/加载命令(REQ-259)
1 parent 8cab673 commit c57ce0c

9 files changed

Lines changed: 266 additions & 8 deletions

File tree

‎app/src-tauri/Cargo.lock‎

Lines changed: 19 additions & 3 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎app/src-tauri/Cargo.toml‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,8 @@ sherpa-onnx = { version = "1.13", default-features = false, features = ["shared"
7272
# .cargo/config.toml 注入);仅取 ndarray 输入通道,禁 download-binaries
7373
# (crate 本地缓存 + config 零网络下载;模型文件另走 model_registry 分发)
7474
ort = { version = "2.0.0-rc.13", default-features = false, features = ["ndarray"] }
75+
# ort ndarray 通道需要直接依赖同版本 ndarray(构建张量输入)
76+
ndarray = "0.16"
7577
# v0.12.0 构建修复:sherpa-onnx 由 static(MT 静态库,内嵌 onnxruntime 符号)切 shared
7678
# (DLL 动态链接)——否则与 oar-ocr 的 ort-sys(MD 动态链接 onnxruntime.dll)链接时
7779
# 发生 LNK2038 RuntimeLibrary 冲突 + LNK2005 onnxruntime 符号重复;shared 包

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

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -232,6 +232,8 @@ pub fn setup_app_state(app: &mut tauri::App) -> Result<(), String> {
232232
db,
233233
engines,
234234
streaming_models,
235+
// REQ-259(v0.19.5):kb embedding 引擎槽(Noop 默认——FTS-only 降级)
236+
embedding_slot: std::sync::Arc::new(std::sync::Mutex::new(Default::default())),
235237
#[cfg(target_os = "windows")]
236238
live_session: LiveSessionManager::new(),
237239
model_downloader: ModelDownloader::new(),

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

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -88,6 +88,9 @@ pub struct AppState {
8888
pub model_downloader: ModelDownloader,
8989
/// 应用句柄(事件推送 live:* / model:*)
9090
pub app: tauri::AppHandle,
91+
/// REQ-259(v0.19.5):kb 语义 embedding 引擎槽(Noop 默认= FTS-only 降级;
92+
/// 模型下载/就绪后经 kb_embedding_load 换入 Onnx——锁内 read-modify-write)
93+
pub embedding_slot: std::sync::Arc<std::sync::Mutex<crate::kb_embed::EmbeddingSlot>>,
9194
/// v0.18.2(REQ-251):目标规划并发互斥——防多窗口/双击重复扣费。
9295
/// 同步规划调用无任务去重表(ai_refine_start 的按会话去重先例);
9396
/// Arc<AtomicBool>(AppState Clone 传播)+ swap 占位(async command

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

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ use tauri::ipc::Channel;
1212
use tauri::State;
1313

1414
use crate::commands::AppState;
15+
use crate::kb_embed::EmbeddingEngine;
1516
use crate::kb_reindex::{KbIndexStats, KbReindexReport};
1617
use crate::kb_search::{KbHit, KB_SEARCH_DEFAULT_LIMIT, KB_SEARCH_MAX_LIMIT};
1718

@@ -56,6 +57,70 @@ pub fn kb_index_stats(state: State<'_, AppState>) -> Result<KbIndexStats, String
5657
state.db.kb_index_stats().map_err(|e| e.to_string())
5758
}
5859

60+
/// kb embedding 模型相对目录(model_dir 下;与 speaker 下载器同约定)
61+
const EMBEDDING_MODEL_REL: &str = "embedding/bge-small-zh-v1.5";
62+
63+
/// 引擎状态视图(设置页「学习库」段数据源——无模型如实显示 noop 与原因)。
64+
#[derive(Serialize)]
65+
#[serde(rename_all = "camelCase")]
66+
pub struct EmbeddingStatusView {
67+
/// noop | onnx
68+
pub kind: String,
69+
/// 引擎是否可用(dim 已知 = 可推理)
70+
pub ready: bool,
71+
/// 输出维度(不可用 None)
72+
pub dim: Option<usize>,
73+
/// 模型目录(诊断)
74+
pub model_dir: String,
75+
/// 当前状态细节(noop=未配置;onnx=就绪/最近失败原因保留由加载命令报错)
76+
pub detail: String,
77+
}
78+
79+
/// 读取 embedding 引擎状态(只读;锁内瞬时快照)。
80+
#[tauri::command]
81+
pub fn kb_embedding_status(state: State<'_, AppState>) -> Result<EmbeddingStatusView, String> {
82+
let slot = state
83+
.embedding_slot
84+
.lock()
85+
.map_err(|_| "embedding 引擎锁中毒".to_string())?;
86+
let dim = slot.engine.dims();
87+
Ok(EmbeddingStatusView {
88+
kind: slot.kind.to_string(),
89+
ready: dim.is_some(),
90+
dim,
91+
model_dir: state.model_dir.join(EMBEDDING_MODEL_REL).to_string_lossy().into_owned(),
92+
detail: if dim.is_some() { "本地模型就绪".to_string() } else { "未配置本地模型(检索按 FTS-only 精度工作)".to_string() },
93+
})
94+
}
95+
96+
/// 加载(或重载)本地 ONNX embedding 引擎。
97+
///
98+
/// @ai-context: 模型文件由下载命令/分发先落位(models/embedding/bge-small-zh-
99+
/// v1.5/{model_quantized.onnx,vocab.txt});本命令只做加载与换槽:
100+
/// 成功 → 槽位切 onnx(后续 reindex 按新引擎重嵌);失败 → 槽位
101+
/// 保持原样并如实报错(不静默降级——状态命令仍显示旧态)。
102+
#[tauri::command]
103+
pub fn kb_embedding_load(state: State<'_, AppState>) -> Result<EmbeddingStatusView, String> {
104+
let dir = state.model_dir.join(EMBEDDING_MODEL_REL);
105+
let engine = crate::kb_embed_onnx::OnnxEmbedding::try_load(&dir)?;
106+
let mut slot = state
107+
.embedding_slot
108+
.lock()
109+
.map_err(|_| "embedding 引擎锁中毒".to_string())?;
110+
let dim = engine.dims();
111+
*slot = crate::kb_embed::EmbeddingSlot {
112+
engine: Box::new(engine),
113+
kind: "onnx",
114+
};
115+
Ok(EmbeddingStatusView {
116+
kind: "onnx".to_string(),
117+
ready: true,
118+
dim,
119+
model_dir: dir.to_string_lossy().into_owned(),
120+
detail: format!("本地模型就绪(dim={})——请在「学习库」段重建索引以回填向量", dim.unwrap_or(0)),
121+
})
122+
}
123+
59124
/// 全量重建(派生索引兜底闸——后台任务 + 进度事件;成功/失败逐源如实报告)。
60125
#[tauri::command]
61126
pub async fn kb_reindex_all(

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

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,8 +8,9 @@
88
//! @ai-context: Onnx/Ollama 具体引擎后续模块实现(OnnxEmbedding 需 ort + 模型
99
//! 文件 + BERT 分词——模型分发复用 model_registry);本模块红线:
1010
//! 引擎产物只是派生索引材料,绝不写结构层(人工裁决闸门铁律)。
11-
//! 注意:引擎接线(kb_search 混合 RRF / reindex 向量回填)落地前,公共 API
12-
//! 尚未被生产路径引用——dead_code 临时豁免,接线轮必须移除本属性(TODO REQ-259)。
11+
//! 注意:向量编解码与 cosine top-K 的检索合流接线(kb_search 混合 RRF /
12+
//! reindex 回填)落地前尚未被生产路径引用——dead_code 临时豁免,混合检索
13+
//! 接线轮必须移除本属性(TODO REQ-259);引擎槽/契约/Noop 已被命令层引用。
1314
#![allow(dead_code)]
1415

1516
use std::error::Error;
@@ -45,6 +46,20 @@ impl EmbeddingEngine for NoopEmbedding {
4546
}
4647
}
4748

49+
/// 引擎槽(AppState 持有;状态命令读、加载命令换入 Onnx——锁内
50+
/// read-modify-write,与词表/开关同模式防 TOCTOU)。
51+
pub struct EmbeddingSlot {
52+
pub engine: Box<dyn EmbeddingEngine>,
53+
/// 引擎标识(noop | onnx——状态命令如实上报)
54+
pub kind: &'static str,
55+
}
56+
57+
impl Default for EmbeddingSlot {
58+
fn default() -> Self {
59+
Self { engine: Box::new(NoopEmbedding), kind: "noop" }
60+
}
61+
}
62+
4863
/// 向量 → f32le BLOB(与 kb_chunks.embedding 列一一对应;长度=dim×4)
4964
pub fn encode_embedding(vec: &[f32]) -> Vec<u8> {
5065
vec.iter().flat_map(|v| v.to_le_bytes()).collect()

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

Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,154 @@
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+
}

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

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,9 +11,6 @@
1111
//! 同分发);加载失败 → 引擎不可用(外层诚实报错,检索自动降级)。
1212
//! @ai-context: 特殊标记取自词表查找(防御换名);词表缺 [UNK]/[CLS]/[SEP]
1313
//! → 加载失败(模型不完整,宁缺勿错)。
14-
//! 注意:引擎模块(kb_embed_onnx)接线前本模块无生产调用——dead_code 临时
15-
//! 豁免,引擎接线轮必须移除本属性(TODO REQ-259)。
16-
#![allow(dead_code)]
1714
1815
use std::collections::HashMap;
1916
use std::path::Path;

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

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -266,6 +266,8 @@ mod kb_discovery;
266266
mod kb_embed;
267267
// REQ-259(v0.19.5):bge-small-zh BERT WordPiece 分词(vocab.txt 加载/编码)
268268
mod kb_embed_tokenizer;
269+
// REQ-259(v0.19.5):bge-small-zh ONNX 推理引擎(ort 封装 + CLS/L2)
270+
mod kb_embed_onnx;
269271
mod commands_kb_discovery;
270272
mod concept_weakness;
271273
mod goal_interview;
@@ -597,6 +599,8 @@ pub fn run() {
597599
commands_kb::kb_search,
598600
commands_kb::kb_index_stats,
599601
commands_kb::kb_reindex_all,
602+
commands_kb::kb_embedding_status,
603+
commands_kb::kb_embedding_load,
600604
// v0.19.1(REQ-260):学习库问答生成开关与预算档位(设置段读写)
601605
commands_ai_settings::ai_set_kb_qa,
602606
// v0.19.3(REQ-261):检索建议(发现路径——默认关,建议制)

0 commit comments

Comments
 (0)