memoryweave/rust/src/quality_backtrace.rs

217 lines
6.4 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// 蒸馏质量回溯 — LLM 反向验证 L0 → L1 蒸馏质量
// 深度整合步骤 4采样 20 条 → L1→L0' 重建 → cosine 比较
use reqwest::blocking::Client;
use serde::{Deserialize, Serialize};
use std::time::Duration;
/// 质量回溯样本
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QualitySample {
pub memory_id: String,
pub distilled_content: String, // L1 蒸馏后内容
pub original_episode: String, // L0 原始对话
pub category: String,
pub tier: String, // fresh / core / useful / not-useful
}
/// 回溯结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QualityBacktraceResult {
pub score: f64, // 整体 cos(L0, L0') 平均
pub low_info_loss: usize, // score < 0.7 的条目数
pub hallucinations: usize, // 疑似幻觉数
pub total: usize,
pub details: Vec<SampleBacktrace>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SampleBacktrace {
pub memory_id: String,
pub cosine: f64,
pub verdict: String, // pass / retry / hallucination
pub re_distilled: String,
}
/// 质量回溯器
pub struct QualityBacktracer {
pub llm_endpoint: String, // LLM API 端点 (new-api / molfang)
pub llm_model: String,
pub client: Client,
pub score_threshold: f64, // 信息损失阈值 (默认 0.7)
}
impl QualityBacktracer {
pub fn new(llm_endpoint: &str, llm_model: &str) -> Self {
Self {
llm_endpoint: llm_endpoint.to_string(),
llm_model: llm_model.to_string(),
client: Client::builder()
.timeout(Duration::from_secs(30))
.build()
.unwrap_or_default(),
score_threshold: 0.7,
}
}
/// 执行质量回溯
pub fn backtrace(
&self,
samples: &[QualitySample],
embed_fn: &dyn Fn(&str) -> Result<Vec<f32>, Box<dyn std::error::Error>>,
) -> Result<QualityBacktraceResult, Box<dyn std::error::Error>> {
let mut details = Vec::new();
let mut low_info_count = 0usize;
let mut hallucination_count = 0usize;
let mut total_cosine = 0.0f64;
for sample in samples {
// Step 1: LLM 反向重建 L0' from L1
let re_distilled = match self.reverse_distill(&sample.distilled_content) {
Ok(r) => r,
Err(e) => {
eprintln!("[quality] 反向蒸馏失败 for {}: {}", sample.memory_id, e);
continue;
}
};
// Step 2: 计算 cos(L0, L0')
let l0_vec = embed_fn(&sample.original_episode)?;
let l0_prime_vec = embed_fn(&re_distilled)?;
let cosine = cosine_similarity(&l0_vec, &l0_prime_vec);
// Step 3: 判定
let verdict = if cosine < 0.5 {
hallucination_count += 1;
"hallucination"
} else if cosine < self.score_threshold {
low_info_count += 1;
"retry"
} else {
"pass"
};
total_cosine += cosine;
details.push(SampleBacktrace {
memory_id: sample.memory_id.clone(),
cosine,
verdict: verdict.to_string(),
re_distilled,
});
}
let total = details.len();
let score = if total > 0 { total_cosine / total as f64 } else { 0.0 };
let result = QualityBacktraceResult {
score,
low_info_loss: low_info_count,
hallucinations: hallucination_count,
total,
details,
};
eprintln!(
"[quality] backtrace {} samples: score={:.3}, low_info={}, hallucinations={}",
total, score, low_info_count, hallucination_count,
);
Ok(result)
}
/// LLM 反向蒸馏 — L1 → L0' 重建
fn reverse_distill(&self, distilled: &str) -> Result<String, Box<dyn std::error::Error>> {
let prompt = format!(
"你看到了一条蒸馏后的记忆事实:\n\n{}\n\n请还原这段事实可能来自怎样的原始对话。只输出还原后的对话文本,不要解释。",
distilled
);
let body = serde_json::json!({
"model": self.llm_model,
"messages": [
{"role": "user", "content": prompt}
],
"temperature": 0.3,
"max_tokens": 200
});
let resp = self.client
.post(&self.llm_endpoint)
.header("Content-Type", "application/json")
.body(serde_json::to_string(&body)?)
.send()?
.text()?;
let parsed: serde_json::Value = serde_json::from_str(&resp)?;
let content = parsed["choices"][0]["message"]["content"]
.as_str()
.unwrap_or("")
.to_string();
Ok(content)
}
}
/// 分层抽样策略:从不同 tier 抽取样本
pub fn stratified_sample(
samples: &[QualitySample],
fresh: usize,
core: usize,
useful: usize,
not_useful: usize,
) -> Vec<QualitySample> {
let mut selected = Vec::new();
let mut by_tier: std::collections::HashMap<&str, Vec<&QualitySample>> =
std::collections::HashMap::new();
for s in samples {
by_tier.entry(&s.tier).or_insert_with(Vec::new).push(s);
}
let tiers = [
("fresh", fresh),
("core", core),
("useful", useful),
("not_useful", not_useful),
];
for (tier, count) in &tiers {
if let Some(list) = by_tier.get(tier) {
let take = (*count).min(list.len());
for s in list.iter().take(take) {
selected.push((*s).clone());
}
}
}
eprintln!(
"[quality] stratified sample: {} total (fresh={}, core={}, useful={}, not_useful={})",
selected.len(), fresh, core, useful, not_useful
);
selected
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> f64 {
let n = a.len().min(b.len());
if n == 0 {
return 0.0;
}
let (mut dot, mut na, mut nb) = (0.0f64, 0.0f64, 0.0f64);
for i in 0..n {
let ai = a[i] as f64;
let bi = b[i] as f64;
dot += ai * bi;
na += ai * ai;
nb += bi * bi;
}
let denom = na.sqrt() * nb.sqrt();
if denom > 1e-10 {
(dot / denom).max(0.0)
} else {
0.0
}
}