1use async_trait::async_trait;
4
5use lc_core::BaseChatModel;
6use lc_embeddings::{cosine_similarity, Embeddings};
7use lc_schema::Message;
8
9use super::criteria::{EvalError, Evaluator, Score};
10
11pub struct ExactMatch;
12
13#[async_trait]
14impl Evaluator for ExactMatch {
15 async fn eval(
16 &self,
17 _input: &str,
18 prediction: &str,
19 reference: &str,
20 ) -> Result<Score, EvalError> {
21 let matched = prediction.trim() == reference.trim();
22 let v = if matched { 1.0 } else { 0.0 };
23 let label = if matched { "match" } else { "mismatch" };
24 Ok(Score::new(v).with_label(label))
25 }
26 fn name(&self) -> &str {
27 "exact_match"
28 }
29}
30
31pub struct StringDistance;
32
33impl StringDistance {
34 fn levenshtein(a: &str, b: &str) -> usize {
35 let a: Vec<char> = a.chars().collect();
36 let b: Vec<char> = b.chars().collect();
37 let (m, n) = (a.len(), b.len());
38 if m == 0 {
39 return n;
40 }
41 if n == 0 {
42 return m;
43 }
44 let mut prev: Vec<usize> = (0..=n).collect();
45 let mut curr: Vec<usize> = vec![0; n + 1];
46 for i in 1..=m {
47 curr[0] = i;
48 for j in 1..=n {
49 let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
50 curr[j] = (prev[j] + 1).min(curr[j - 1] + 1).min(prev[j - 1] + cost);
51 }
52 std::mem::swap(&mut prev, &mut curr);
53 }
54 prev[n]
55 }
56}
57
58#[async_trait]
59impl Evaluator for StringDistance {
60 async fn eval(
61 &self,
62 _input: &str,
63 prediction: &str,
64 reference: &str,
65 ) -> Result<Score, EvalError> {
66 let dist = Self::levenshtein(prediction, reference) as f64;
67 let max_len = prediction.chars().count().max(reference.chars().count()) as f64;
68 let score = if max_len == 0.0 {
69 1.0
70 } else {
71 1.0 - dist / max_len
72 };
73 Ok(Score::new(score))
74 }
75 fn name(&self) -> &str {
76 "string_distance"
77 }
78}
79
80pub struct EmbeddingSimilarity<E: Embeddings> {
81 embeddings: E,
82}
83
84impl<E: Embeddings> EmbeddingSimilarity<E> {
85 pub fn new(embeddings: E) -> Self {
86 Self { embeddings }
87 }
88}
89
90#[async_trait]
91impl<E: Embeddings> Evaluator for EmbeddingSimilarity<E> {
92 async fn eval(
93 &self,
94 _input: &str,
95 prediction: &str,
96 reference: &str,
97 ) -> Result<Score, EvalError> {
98 let p = self
99 .embeddings
100 .embed_query(prediction)
101 .await
102 .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
103 let r = self
104 .embeddings
105 .embed_query(reference)
106 .await
107 .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
108 let sim = cosine_similarity(&p, &r).unwrap_or(0.0);
109 let v = ((sim + 1.0) / 2.0).clamp(0.0, 1.0);
110 Ok(Score::new(v as f64))
111 }
112 fn name(&self) -> &str {
113 "embedding_similarity"
114 }
115}
116
117pub struct LLMAsJudge<M: BaseChatModel> {
118 judge: M,
119 rubric: String,
120 max_score: u8,
121}
122
123const DEFAULT_RUBRIC: &str = "\
124正确性:回答是否事实准确、是否与参考答案的核心意思一致。
125完整性:是否完整回答了输入的问题或指令。
126清晰性:表达是否清晰、无歧义、无冗余。";
127
128impl<M: BaseChatModel> LLMAsJudge<M> {
129 pub fn new(judge: M) -> Self {
130 Self {
131 judge,
132 rubric: DEFAULT_RUBRIC.to_string(),
133 max_score: 10,
134 }
135 }
136 pub fn with_rubric(mut self, rubric: impl Into<String>) -> Self {
137 self.rubric = rubric.into();
138 self
139 }
140 pub fn with_max_score(mut self, max_score: u8) -> Self {
141 self.max_score = max_score.max(1);
142 self
143 }
144
145 fn build_prompt(&self, input: &str, prediction: &str, reference: &str) -> (String, String) {
146 let system = format!(
147 "你是一个严格、公正的评估员。请根据以下评分标准对待评估的回答打分。\n\n评分标准:\n{rubric}\n\n打分范围:0 到 {max}(0 = 完全错误或无关,{max} = 完全正确)。\n\n要求:先在 reason 字段写出简短分析,再在 score 字段给出分数。\n只输出一行 JSON,格式为:{{\"reason\":\"...\",\"score\":N}}",
148 rubric = self.rubric, max = self.max_score
149 );
150 let user =
151 format!(
152 "输入:\n{input}\n\n参考答案:\n{reference}\n\n待评估的回答:\n{prediction}\n\n请评估。",
153 input = input, reference = reference, prediction = prediction,
154 );
155 (system, user)
156 }
157}
158
159#[async_trait]
160impl<M: BaseChatModel> Evaluator for LLMAsJudge<M> {
161 async fn eval(
162 &self,
163 input: &str,
164 prediction: &str,
165 reference: &str,
166 ) -> Result<Score, EvalError> {
167 let (system, user) = self.build_prompt(input, prediction, reference);
168 let result = self
169 .judge
170 .chat_with_system(system, vec![Message::human(user)])
171 .await
172 .map_err(|e| EvalError::PredictorError(e.to_string()))?;
173 let raw = result.content;
174 let value = parse_score(&raw, self.max_score).ok_or_else(|| {
175 EvalError::ParseError(format!("无法从裁判回复解析分数: {}", truncate(&raw, 200)))
176 })?;
177 Ok(Score::new(value).with_label("llm_judge"))
178 }
179 fn name(&self) -> &str {
180 "llm_as_judge"
181 }
182}
183
184fn parse_score(raw: &str, max_score: u8) -> Option<f64> {
185 let max = max_score as f64;
186 let n = extract_json_score(raw)
187 .or_else(|| find_number_after_keyword(raw, "score"))
188 .or_else(|| find_number_after_keyword(raw, "分数"))
189 .or_else(|| first_number(raw))?;
190 Some((n / max).clamp(0.0, 1.0))
191}
192
193fn extract_json_score(raw: &str) -> Option<f64> {
194 let start = raw.find('{')?;
195 let end = raw.rfind('}')?;
196 if end < start {
197 return None;
198 }
199 let val: serde_json::Value = serde_json::from_str(&raw[start..=end]).ok()?;
200 val.get("score")?.as_f64()
201}
202
203fn find_number_after_keyword(raw: &str, keyword: &str) -> Option<f64> {
204 let lower_raw = raw.to_lowercase();
205 let lower_kw = keyword.to_lowercase();
206 let idx = lower_raw.find(lower_kw.as_str())?;
207 first_number(&raw[idx + lower_kw.len()..])
208}
209
210fn first_number(s: &str) -> Option<f64> {
211 let mut buf = String::new();
212 let mut started = false;
213 for c in s.chars() {
214 if c.is_ascii_digit() || c == '.' {
215 started = true;
216 buf.push(c);
217 } else if started {
218 break;
219 }
220 }
221 if buf.is_empty() {
222 return None;
223 }
224 buf.parse::<f64>().ok().or_else(|| {
225 let int_part: String = buf.chars().take_while(|c| c.is_ascii_digit()).collect();
226 int_part.parse::<f64>().ok()
227 })
228}
229
230fn truncate(s: &str, max: usize) -> String {
231 if s.chars().count() <= max {
232 s.to_string()
233 } else {
234 let truncated: String = s.chars().take(max).collect();
235 format!("{}...", truncated)
236 }
237}
238
239#[cfg(test)]
240mod tests {
241 use super::*;
242 use lc_embeddings::MockEmbeddings;
243
244 #[tokio::test]
245 async fn test_exact_match() {
246 let ev = ExactMatch;
247 assert_eq!(ev.eval("", "hello", "hello").await.unwrap().value, 1.0);
248 assert_eq!(ev.eval("", "hello", "world").await.unwrap().value, 0.0);
249 assert_eq!(ev.eval("", " yes ", "yes").await.unwrap().value, 1.0);
250 }
251
252 #[tokio::test]
253 async fn test_string_distance() {
254 let ev = StringDistance;
255 let s = ev.eval("", "hello", "hello").await.unwrap().value;
256 assert!((s - 1.0).abs() < 1e-9);
257 let s = ev.eval("", "kitten", "sitting").await.unwrap().value;
258 assert!((s - (1.0 - 3.0 / 7.0)).abs() < 1e-6);
259 }
260
261 #[tokio::test]
262 async fn test_embedding_similarity_identical() {
263 let ev = EmbeddingSimilarity::new(MockEmbeddings::new(32));
264 let s = ev.eval("", "hello", "hello").await.unwrap().value;
265 assert!((s - 1.0).abs() < 1e-6);
266 }
267
268 #[test]
269 fn test_parse_score_unit() {
270 assert!((parse_score(r#"{"score":10}"#, 10).unwrap() - 1.0).abs() < 1e-9);
271 assert!((parse_score(r#"{"score":7.5}"#, 10).unwrap() - 0.75).abs() < 1e-9);
272 assert!((parse_score(r#"{"score":12}"#, 10).unwrap() - 1.0).abs() < 1e-9);
273 assert!(parse_score("no number here", 10).is_none());
274 }
275}