Skip to main content

lc_evaluation/
results.rs

1//! Built-in evaluators: ExactMatch, StringDistance, EmbeddingSimilarity, LLMAsJudge.
2
3use async_trait::async_trait;
4use serde::Deserialize;
5
6use lc_core::tools::ToolDefinition;
7use lc_core::BaseChatModel;
8use lc_embeddings::{cosine_similarity, Embeddings};
9use lc_schema::Message;
10
11use lc_core::judge::{structured_call, truncate, StructuredJudgeError};
12
13use super::criteria::{EvalError, Evaluator, Score};
14
15/// 精确匹配评测器:预测与参考答案 trim 后完全一致得 1.0,否则 0.0。
16pub struct ExactMatch;
17
18#[async_trait]
19impl Evaluator for ExactMatch {
20    async fn eval(
21        &self,
22        _input: &str,
23        prediction: &str,
24        reference: &str,
25    ) -> Result<Score, EvalError> {
26        let matched = prediction.trim() == reference.trim();
27        let v = if matched { 1.0 } else { 0.0 };
28        let label = if matched { "match" } else { "mismatch" };
29        Ok(Score::new(v).with_label(label))
30    }
31    fn name(&self) -> &str {
32        "exact_match"
33    }
34}
35
36/// 字符串距离评测器:基于 Levenshtein 编辑距离归一化打分。
37pub struct StringDistance;
38
39impl StringDistance {
40    fn levenshtein(a: &str, b: &str) -> usize {
41        let a: Vec<char> = a.chars().collect();
42        let b: Vec<char> = b.chars().collect();
43        let (m, n) = (a.len(), b.len());
44        if m == 0 {
45            return n;
46        }
47        if n == 0 {
48            return m;
49        }
50        let mut prev: Vec<usize> = (0..=n).collect();
51        let mut curr: Vec<usize> = vec![0; n + 1];
52        for i in 1..=m {
53            curr[0] = i;
54            for j in 1..=n {
55                let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
56                curr[j] = (prev[j] + 1).min(curr[j - 1] + 1).min(prev[j - 1] + cost);
57            }
58            std::mem::swap(&mut prev, &mut curr);
59        }
60        prev[n]
61    }
62}
63
64#[async_trait]
65impl Evaluator for StringDistance {
66    async fn eval(
67        &self,
68        _input: &str,
69        prediction: &str,
70        reference: &str,
71    ) -> Result<Score, EvalError> {
72        let dist = Self::levenshtein(prediction, reference) as f64;
73        let max_len = prediction.chars().count().max(reference.chars().count()) as f64;
74        let score = if max_len == 0.0 {
75            1.0
76        } else {
77            1.0 - dist / max_len
78        };
79        Ok(Score::new(score))
80    }
81    fn name(&self) -> &str {
82        "string_distance"
83    }
84}
85
86/// 嵌入相似度评测器:用预测与参考答案的嵌入向量余弦相似度打分。
87pub struct EmbeddingSimilarity<E: Embeddings> {
88    embeddings: E,
89}
90
91impl<E: Embeddings> EmbeddingSimilarity<E> {
92    /// 创建嵌入相似度评测器。
93    pub fn new(embeddings: E) -> Self {
94        Self { embeddings }
95    }
96}
97
98#[async_trait]
99impl<E: Embeddings> Evaluator for EmbeddingSimilarity<E> {
100    async fn eval(
101        &self,
102        _input: &str,
103        prediction: &str,
104        reference: &str,
105    ) -> Result<Score, EvalError> {
106        let p = self
107            .embeddings
108            .embed_query(prediction)
109            .await
110            .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
111        let r = self
112            .embeddings
113            .embed_query(reference)
114            .await
115            .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
116        // P2-7: cosine 长度不匹配是数据缺陷(向量维度不一致),不再吞成 0.0 静默降级
117        let sim =
118            cosine_similarity(&p, &r).map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
119        let v = ((sim + 1.0) / 2.0).clamp(0.0, 1.0);
120        Ok(Score::new(v as f64))
121    }
122    fn name(&self) -> &str {
123        "embedding_similarity"
124    }
125}
126
127/// LLM 裁判评测器:让 LLM 按评分标准对预测打分(0 到 `max_score`)。
128pub struct LLMAsJudge<M: BaseChatModel> {
129    judge: M,
130    rubric: String,
131    max_score: u8,
132}
133
134const DEFAULT_RUBRIC: &str = "\
135正确性:回答是否事实准确、是否与参考答案的核心意思一致。
136完整性:是否完整回答了输入的问题或指令。
137清晰性:表达是否清晰、无歧义、无冗余。";
138
139impl<M: BaseChatModel> LLMAsJudge<M> {
140    /// 创建使用默认评分标准与 10 分制的 LLM 裁判评测器。
141    pub fn new(judge: M) -> Self {
142        Self {
143            judge,
144            rubric: DEFAULT_RUBRIC.to_string(),
145            max_score: 10,
146        }
147    }
148    /// 设置自定义评分标准(builder 风格)。
149    pub fn with_rubric(mut self, rubric: impl Into<String>) -> Self {
150        self.rubric = rubric.into();
151        self
152    }
153    /// 设置最高分(最小为 1)。
154    pub fn with_max_score(mut self, max_score: u8) -> Self {
155        self.max_score = max_score.max(1);
156        self
157    }
158
159    fn build_prompt(&self, input: &str, prediction: &str, reference: &str) -> (String, String) {
160        let system = format!(
161            "你是一个严格、公正的评估员。请根据以下评分标准对待评估的回答打分。\n\n评分标准:\n{rubric}\n\n打分范围:0 到 {max}(0 = 完全错误或无关,{max} = 完全正确)。\n\n要求:先在 reason 字段写出简短分析,再在 score 字段给出分数。\n只输出一行 JSON,格式为:{{\"reason\":\"...\",\"score\":N}}",
162            rubric = self.rubric, max = self.max_score
163        );
164        let user =
165            format!(
166            "输入:\n{input}\n\n参考答案:\n{reference}\n\n待评估的回答:\n{prediction}\n\n请评估。",
167            input = input, reference = reference, prediction = prediction,
168        );
169        (system, user)
170    }
171}
172
173/// 结构化评分参数(经 tool_calls 返回)。
174#[derive(Debug, Deserialize)]
175struct ScoreArgs {
176    score: f64,
177    /// 让 LLM 附上简短理由(改善打分质量),当前不消费。
178    #[serde(default)]
179    #[allow(dead_code)]
180    reason: String,
181}
182
183/// 构建评分工具:让 LLM 以 `{"score": 0..max, "reason": "..."}` 提交评分。
184fn score_tool(max_score: u8) -> ToolDefinition {
185    ToolDefinition::new(
186        "submit_evaluation",
187        "提交你对回答的评分。score 为 0 到 max 的整数,reason 给出简短分析。",
188    )
189    .with_parameters(serde_json::json!({
190        "type": "object",
191        "properties": {
192            "score": {
193                "type": "integer",
194                "minimum": 0,
195                "maximum": max_score,
196                "description": "0 到 max 的整数分数"
197            },
198            "reason": { "type": "string", "description": "简短分析" }
199        },
200        "required": ["score", "reason"]
201    }))
202}
203
204#[async_trait]
205impl<M: BaseChatModel> Evaluator for LLMAsJudge<M> {
206    async fn eval(
207        &self,
208        input: &str,
209        prediction: &str,
210        reference: &str,
211    ) -> Result<Score, EvalError> {
212        let (system, user) = self.build_prompt(input, prediction, reference);
213        let messages = vec![Message::system(system), Message::human(user)];
214
215        // P0-1: 优先结构化输出(tool_calls);不支持工具绑定的模型走文本解析回落。
216        let args: ScoreArgs =
217            structured_call(&self.judge, score_tool(self.max_score), messages, |raw| {
218                let norm = parse_score(raw, self.max_score).ok_or_else(|| {
219                    StructuredJudgeError::Parse(format!(
220                        "failed to parse score from judge reply: {}",
221                        truncate(raw, 200)
222                    ))
223                })?;
224                Ok(ScoreArgs {
225                    score: norm * self.max_score as f64,
226                    reason: String::new(),
227                })
228            })
229            .await?;
230
231        let value = (args.score / self.max_score as f64).clamp(0.0, 1.0);
232        Ok(Score::new(value).with_label("llm_judge"))
233    }
234    fn name(&self) -> &str {
235        "llm_as_judge"
236    }
237}
238
239fn parse_score(raw: &str, max_score: u8) -> Option<f64> {
240    let max = max_score as f64;
241    let n = extract_json_score(raw)
242        .or_else(|| find_number_after_keyword(raw, "score"))
243        .or_else(|| find_number_after_keyword(raw, "分数"))
244        .or_else(|| first_number(raw))?;
245    Some((n / max).clamp(0.0, 1.0))
246}
247
248fn extract_json_score(raw: &str) -> Option<f64> {
249    let start = raw.find('{')?;
250    let end = raw.rfind('}')?;
251    if end < start {
252        return None;
253    }
254    let val: serde_json::Value = serde_json::from_str(&raw[start..=end]).ok()?;
255    val.get("score")?.as_f64()
256}
257
258fn find_number_after_keyword(raw: &str, keyword: &str) -> Option<f64> {
259    let lower_raw = raw.to_lowercase();
260    let lower_kw = keyword.to_lowercase();
261    let idx = lower_raw.find(lower_kw.as_str())?;
262    first_number(&raw[idx + lower_kw.len()..])
263}
264
265fn first_number(s: &str) -> Option<f64> {
266    let mut buf = String::new();
267    let mut started = false;
268    for c in s.chars() {
269        if c.is_ascii_digit() || c == '.' {
270            started = true;
271            buf.push(c);
272        } else if started {
273            break;
274        }
275    }
276    if buf.is_empty() {
277        return None;
278    }
279    buf.parse::<f64>().ok().or_else(|| {
280        let int_part: String = buf.chars().take_while(|c| c.is_ascii_digit()).collect();
281        int_part.parse::<f64>().ok()
282    })
283}
284
285#[cfg(test)]
286mod tests {
287    use super::*;
288    use lc_embeddings::{EmbeddingError, MockEmbeddings};
289
290    /// 让 embed_query 按文本返回不同维度的向量,复现"向量长度不匹配"。
291    struct MismatchedDimEmbeddings;
292
293    #[async_trait]
294    impl Embeddings for MismatchedDimEmbeddings {
295        async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
296            if text == "pred" {
297                Ok(vec![1.0; 4])
298            } else {
299                Ok(vec![1.0; 8])
300            }
301        }
302        fn dimension(&self) -> usize {
303            4
304        }
305        fn model_name(&self) -> &str {
306            "mismatched-dim"
307        }
308    }
309
310    #[tokio::test]
311    async fn test_exact_match() {
312        let ev = ExactMatch;
313        assert_eq!(ev.eval("", "hello", "hello").await.unwrap().value, 1.0);
314        assert_eq!(ev.eval("", "hello", "world").await.unwrap().value, 0.0);
315        assert_eq!(ev.eval("", "  yes  ", "yes").await.unwrap().value, 1.0);
316    }
317
318    #[tokio::test]
319    async fn test_string_distance() {
320        let ev = StringDistance;
321        let s = ev.eval("", "hello", "hello").await.unwrap().value;
322        assert!((s - 1.0).abs() < 1e-9);
323        let s = ev.eval("", "kitten", "sitting").await.unwrap().value;
324        assert!((s - (1.0 - 3.0 / 7.0)).abs() < 1e-6);
325    }
326
327    #[tokio::test]
328    async fn test_embedding_similarity_identical() {
329        let ev = EmbeddingSimilarity::new(MockEmbeddings::new(32));
330        let s = ev.eval("", "hello", "hello").await.unwrap().value;
331        assert!((s - 1.0).abs() < 1e-6);
332    }
333
334    /// P2-7: 向量长度不匹配是数据缺陷,报错而不是静默返回 0.0。
335    #[tokio::test]
336    async fn test_embedding_similarity_mismatched_dim_errors() {
337        let ev = EmbeddingSimilarity::new(MismatchedDimEmbeddings);
338        let err = ev.eval("", "pred", "ref").await.unwrap_err();
339        assert!(matches!(err, EvalError::EmbeddingError(_)));
340    }
341
342    #[test]
343    fn test_parse_score_unit() {
344        assert!((parse_score(r#"{"score":10}"#, 10).unwrap() - 1.0).abs() < 1e-9);
345        assert!((parse_score(r#"{"score":7.5}"#, 10).unwrap() - 0.75).abs() < 1e-9);
346        assert!((parse_score(r#"{"score":12}"#, 10).unwrap() - 1.0).abs() < 1e-9);
347        assert!(parse_score("no number here", 10).is_none());
348    }
349}