1use serde::{Deserialize, Serialize};
4
5#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
6use crate::backend::CandleBackend;
7#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
8use crate::tasks::generate;
9use crate::InferenceError;
10
11#[derive(Debug, Clone, Serialize, Deserialize)]
13pub struct ClassifyRequest {
14 pub text: String,
16 pub labels: Vec<String>,
18 pub model: Option<String>,
20 #[serde(skip)]
22 pub work_context: Option<car_auth::context::CredentialContext>,
23}
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
27pub struct ClassifyResult {
28 pub label: String,
29 pub score: f64,
30}
31
32#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct OptionProbabilities {
35 pub probabilities: Vec<f64>,
37 pub mass: f64,
40 pub method: String,
42}
43
44#[cfg(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx)))]
45fn log_prob(logits: &[f32], token: u32) -> Result<f64, InferenceError> {
46 let at = logits.get(token as usize).ok_or_else(|| {
47 InferenceError::InferenceFailed(format!(
48 "token {token} is outside the model's {}-entry vocabulary",
49 logits.len()
50 ))
51 })?;
52 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
53 let sum: f64 = logits.iter().map(|&l| ((l as f64) - max).exp()).sum();
54 Ok((*at as f64) - max - sum.ln())
55}
56
57#[cfg(any(
58 test,
59 all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))
60))]
61fn log_sum_exp(values: &[f64]) -> f64 {
62 let max = values.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
63 if max == f64::NEG_INFINITY {
64 return max;
65 }
66 max + values.iter().map(|v| (v - max).exp()).sum::<f64>().ln()
67}
68
69#[cfg(any(
72 test,
73 all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))
74))]
75fn spellings(option: &str) -> Vec<String> {
76 let mut chars = option.chars();
77 let capitalized = match chars.next() {
78 Some(c) => c.to_uppercase().collect::<String>() + chars.as_str(),
79 None => String::new(),
80 };
81 let mut out = vec![option.to_string()];
82 if capitalized != option {
83 out.push(capitalized);
84 }
85 out
86}
87
88#[cfg(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx)))]
91pub fn score_options(
92 backend: &mut crate::backend::SwiftLmBackend,
93 formatted: &str,
94 options: &[String],
95) -> Result<OptionProbabilities, InferenceError> {
96 let prompt = backend.encode(formatted)?;
97 let variants: Vec<Vec<Vec<u32>>> = options
100 .iter()
101 .map(|o| {
102 spellings(o)
103 .iter()
104 .map(|v| backend.tokenize_raw(v))
105 .collect::<Result<Vec<_>, _>>()
106 })
107 .collect::<Result<_, _>>()?;
108 if variants.iter().flatten().any(|t| t.is_empty()) {
109 return Err(InferenceError::InvalidClassifyLabels(
110 "an option encodes to no tokens".into(),
111 ));
112 }
113 let firsts: Vec<std::collections::HashSet<u32>> = variants
114 .iter()
115 .map(|v| v.iter().map(|t| t[0]).collect())
116 .collect();
117 let collide = firsts
118 .iter()
119 .enumerate()
120 .any(|(i, a)| firsts.iter().skip(i + 1).any(|b| !a.is_disjoint(b)));
121 let (scores, method) = if !collide {
122 backend.clear_kv_cache();
123 let logits = backend.forward(&prompt, 0)?;
124 let scores = firsts
125 .iter()
126 .map(|ids| {
127 let lps = ids
128 .iter()
129 .map(|&id| log_prob(&logits, id))
130 .collect::<Result<Vec<_>, _>>()?;
131 Ok(log_sum_exp(&lps))
132 })
133 .collect::<Result<Vec<_>, InferenceError>>()?;
134 (scores, "first_token")
135 } else {
136 let end = backend.token_id("<|im_end|>");
137 let mut scores = Vec::with_capacity(options.len());
138 for option in &variants {
139 let mut per_spelling = Vec::with_capacity(option.len());
140 for tokens in option {
141 let sequence: Vec<u32> = tokens.iter().copied().chain(end).collect();
142 backend.clear_kv_cache();
143 let mut logits = backend.forward(&prompt, 0)?;
144 let mut total = 0.0;
145 for (i, &token) in sequence.iter().enumerate() {
146 total += log_prob(&logits, token)?;
147 if i + 1 < sequence.len() {
148 logits = backend.forward(&[token], prompt.len() + i)?;
149 }
150 }
151 per_spelling.push(total);
152 }
153 scores.push(log_sum_exp(&per_spelling));
154 }
155 (scores, "sequence")
156 };
157 let mass: f64 = scores.iter().map(|s| s.exp()).sum();
158 if mass.is_nan() || mass <= 0.0 {
159 return Err(InferenceError::InferenceFailed(
160 "the model gave none of the options any probability".into(),
161 ));
162 }
163 let top = log_sum_exp(&scores);
164 Ok(OptionProbabilities {
165 probabilities: scores.iter().map(|s| (s - top).exp()).collect(),
166 mass,
167 method: method.into(),
168 })
169}
170
171fn normalize(text: &str) -> String {
174 text.to_lowercase()
175 .split(|c: char| !c.is_alphanumeric())
176 .filter(|w| !w.is_empty())
177 .collect::<Vec<_>>()
178 .join(" ")
179}
180
181const STOPWORDS: &[&str] = &[
184 "a", "an", "the", "i", "is", "are", "am", "be", "of", "to", "in", "on", "it", "this", "that",
185 "and", "or", "for", "with", "my", "me", "you", "your", "can", "t", "do", "not",
186];
187
188fn unspaced_script(text: &str) -> bool {
191 text.chars().any(|c| {
192 matches!(c as u32,
193 0x3040..=0x30FF | 0x3400..=0x4DBF | 0x4E00..=0x9FFF | 0x0E00..=0x0E7F | 0x0E80..=0x0EFF | 0x1780..=0x17FF | 0x1000..=0x109F )
201 })
202}
203
204pub fn validate_labels(labels: &[String]) -> Result<(), InferenceError> {
206 if labels.is_empty() {
207 return Err(InferenceError::InvalidClassifyLabels(
208 "no labels given".into(),
209 ));
210 }
211 let mut seen = std::collections::HashMap::new();
212 for label in labels {
213 let norm = normalize(label);
214 if norm.is_empty() {
215 return Err(InferenceError::InvalidClassifyLabels(format!(
216 "{label:?} has no letters or digits"
217 )));
218 }
219 if let Some(other) = seen.insert(norm, label) {
220 return Err(InferenceError::InvalidClassifyLabels(format!(
221 "{other:?} and {label:?} differ only in case or separators"
222 )));
223 }
224 }
225 Ok(())
226}
227
228pub fn score_reply(reply: &str, labels: &[String]) -> Result<Vec<ClassifyResult>, InferenceError> {
241 validate_labels(labels)?;
242 let no_answer = |reason: &str| InferenceError::ClassifyNoAnswer {
243 reply: reply.trim().chars().take(120).collect(),
244 reason: reason.into(),
245 };
246 let reply_norm = normalize(reply);
247 if let Ok(n) = reply_norm.parse::<usize>() {
248 if (1..=labels.len()).contains(&n) {
249 return Ok(labels
250 .iter()
251 .enumerate()
252 .map(|(i, label)| ClassifyResult {
253 label: label.clone(),
254 score: if i + 1 == n { 1.0 } else { 0.0 },
255 })
256 .collect::<Vec<_>>())
257 .map(|mut results: Vec<ClassifyResult>| {
258 results.sort_by(|a, b| {
259 b.score
260 .partial_cmp(&a.score)
261 .unwrap_or(std::cmp::Ordering::Equal)
262 });
263 results
264 });
265 }
266 }
267 let reply_words: std::collections::HashSet<&str> = reply_norm.split(' ').collect();
268 let padded = format!(" {reply_norm} ");
269 let mut results: Vec<ClassifyResult> = labels
270 .iter()
271 .map(|label| {
272 let label_norm = normalize(label);
273 let contained = padded.contains(&format!(" {label_norm} "))
274 || (unspaced_script(&label_norm) && reply_norm.contains(&label_norm));
275 let score = if reply_norm == label_norm {
276 1.0
277 } else if contained {
278 0.8
279 } else {
280 let content: Vec<&str> = label_norm
281 .split(' ')
282 .filter(|w| !STOPWORDS.contains(w))
283 .collect();
284 let hits = content.iter().filter(|w| reply_words.contains(*w)).count();
285 let share = if content.is_empty() {
286 0.0
287 } else {
288 hits as f64 / content.len() as f64
289 };
290 if share >= 0.5 {
291 0.5 * share
292 } else {
293 0.0
294 }
295 };
296 ClassifyResult {
297 label: label.clone(),
298 score,
299 }
300 })
301 .collect();
302 results.sort_by(|a, b| {
303 b.score
304 .partial_cmp(&a.score)
305 .unwrap_or(std::cmp::Ordering::Equal)
306 });
307 let total: f64 = results.iter().map(|r| r.score).sum();
308 if total <= 0.0 {
309 return Err(no_answer("names none of the labels"));
310 }
311 if results.len() > 1 && results[0].score < 1.0 && results[0].score == results[1].score {
312 return Err(no_answer("names several labels equally"));
313 }
314 for r in &mut results {
315 r.score /= total;
316 }
317 Ok(results)
318}
319
320pub fn system_one_request_body(
325 model: &str,
326 text: &str,
327 labels: &[String],
328) -> Result<String, InferenceError> {
329 validate_labels(labels)?;
330 let criteria = labels
331 .iter()
332 .map(|label| {
333 format!(
334 "{}:{}",
335 serde_json::Value::from(label.as_str()),
336 serde_json::Value::from(normalize(label))
337 )
338 })
339 .collect::<Vec<_>>()
340 .join(",");
341 Ok(format!(
342 "{{\"model\":{},\"state\":{},\"questions\":{{\"label\":{{\"type\":\"choice\",\
343 \"instructions\":{},\"criteria\":{{{criteria}}}}}}}}}",
344 serde_json::Value::from(model),
345 serde_json::json!({ "text": text }),
346 serde_json::Value::from("Classify `text` into one of the labels."),
347 ))
348}
349
350pub fn system_one_results(
354 response: &serde_json::Value,
355 labels: &[String],
356) -> Result<Vec<ClassifyResult>, InferenceError> {
357 let answer = response.pointer("/answers/label").ok_or_else(|| {
358 InferenceError::InferenceFailed(format!("System One response has no answer: {response}"))
359 })?;
360 let choice = answer
361 .get("choice")
362 .and_then(|c| c.as_str())
363 .ok_or_else(|| InferenceError::InferenceFailed("System One answer has no choice".into()))?;
364 if !labels.iter().any(|l| l == choice) {
365 return Err(InferenceError::InferenceFailed(format!(
366 "System One chose {choice:?}, which is not an offered label"
367 )));
368 }
369 let probabilities = answer
370 .get("probabilities")
371 .and_then(|p| p.as_object())
372 .ok_or_else(|| {
373 InferenceError::InferenceFailed("System One answer has no probabilities".into())
374 })?;
375 let mut results = labels
378 .iter()
379 .map(|label| {
380 probabilities
381 .get(label)
382 .and_then(|p| p.as_f64())
383 .map(|score| ClassifyResult {
384 label: label.clone(),
385 score,
386 })
387 .ok_or_else(|| {
388 InferenceError::InferenceFailed(format!(
389 "System One gave no probability for {label:?}"
390 ))
391 })
392 })
393 .collect::<Result<Vec<_>, _>>()?;
394 results.sort_by(|a, b| {
395 b.score
396 .partial_cmp(&a.score)
397 .unwrap_or(std::cmp::Ordering::Equal)
398 });
399 if let Some(i) = results.iter().position(|r| r.label == choice) {
401 let chosen = results.remove(i);
402 results.insert(0, chosen);
403 }
404 Ok(results)
405}
406
407#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
408pub async fn classify(
413 backend: &mut CandleBackend,
414 req: ClassifyRequest,
415) -> Result<Vec<ClassifyResult>, InferenceError> {
416 let labels_str = req
417 .labels
418 .iter()
419 .enumerate()
420 .map(|(i, l)| format!("{}. {}", i + 1, l))
421 .collect::<Vec<_>>()
422 .join("\n");
423
424 let prompt = format!(
425 "Classify the following text into one of these categories:\n\
426 {labels_str}\n\n\
427 Text: {}\n\n\
428 Respond with ONLY the category name, nothing else.",
429 req.text
430 );
431
432 let gen_req = generate::GenerateRequest {
433 work_context: req.work_context.clone(),
434 prompt,
435 model: req.model.clone(),
436 params: generate::GenerateParams {
437 temperature: 0.0, max_tokens: 32,
439 ..Default::default()
440 },
441 context: None,
442 context_stable_prefix: None,
443 tools: None,
444 images: None,
445 messages: None,
446 cache_control: false,
447 response_format: None,
448 intent: None,
449 client_ref: None,
450 expected_row_digest: None,
451 expected_catalog_revision: None,
452 caller: None,
453 };
454
455 let (response, _ttft_ms, _prompt_tokens, _completion_tokens) =
456 generate::generate(backend, gen_req).await?;
457 score_reply(&response, &req.labels)
458 .map_err(|e| InferenceError::InferenceFailed(format!("classify: {e}")))
459}
460
461#[cfg(test)]
462mod score_tests {
463 use super::*;
464
465 #[test]
466 fn classification_context_is_trusted_and_not_serialized() {
467 let context = car_auth::context::CredentialContext {
468 api_base: "https://authority.example".into(),
469 account_id: "original-account".into(),
470 organization_id: Some("parslee".into()),
471 };
472 let request = ClassifyRequest {
473 text: "a task".into(),
474 labels: vec!["work".into()],
475 model: None,
476 work_context: Some(context.clone()),
477 };
478 let mut serialized = serde_json::to_value(&request).unwrap();
479 assert!(serialized.get("work_context").is_none());
480 serialized["work_context"] = serde_json::to_value(context).unwrap();
481 let decoded: ClassifyRequest = serde_json::from_value(serialized).unwrap();
482 assert!(decoded.work_context.is_none());
483 }
484
485 fn labels(names: &[&str]) -> Vec<String> {
486 names.iter().map(|s| s.to_string()).collect()
487 }
488
489 fn no_answer(r: Result<Vec<ClassifyResult>, InferenceError>) -> String {
490 match r {
491 Err(InferenceError::ClassifyNoAnswer { reason, .. }) => reason,
492 other => panic!("expected ClassifyNoAnswer, got {other:?}"),
493 }
494 }
495
496 #[test]
497 fn separators_and_case_do_not_hide_a_label() {
498 let r = score_reply("Out of scope.", &labels(&["transfer", "out_of_scope"])).unwrap();
499 assert_eq!(r[0].label, "out_of_scope");
500 assert!(r[0].score > 0.99);
501 }
502
503 #[test]
504 fn a_label_inside_a_reply_is_matched_on_whole_words() {
505 let r = score_reply(
506 "I think it is transfer money",
507 &labels(&["transfer", "yes"]),
508 )
509 .unwrap();
510 assert_eq!(r[0].label, "transfer");
511 assert_eq!(
512 no_answer(score_reply("yesterday", &labels(&["yes", "no"]))),
513 "names none of the labels"
514 );
515 }
516
517 #[test]
518 fn a_refusal_does_not_match_on_small_words() {
519 let r = score_reply(
520 "I can't classify this",
521 &labels(&["are you a bot", "book hotel"]),
522 );
523 assert_eq!(no_answer(r), "names none of the labels");
524 }
525
526 #[test]
527 fn several_labels_named_equally_is_no_answer() {
528 let r = score_reply(
529 "transfer or balance",
530 &labels(&["transfer", "balance", "timer"]),
531 );
532 assert_eq!(no_answer(r), "names several labels equally");
533 let r = score_reply("hotel", &labels(&["book hotel", "hotel", "hotel reviews"])).unwrap();
535 assert_eq!(r[0].label, "hotel");
536 }
537
538 #[test]
539 fn a_numbered_reply_names_that_label() {
540 let r = score_reply("2.", &labels(&["email", "calendar", "search"])).unwrap();
541 assert_eq!(r[0].label, "calendar");
542 assert_eq!(r[0].score, 1.0);
543 assert_eq!(
544 no_answer(score_reply("7", &labels(&["email", "calendar"]))),
545 "names none of the labels"
546 );
547 }
548
549 #[test]
550 fn an_unspaced_script_label_is_found_inside_the_reply() {
551 let r = score_reply("今天天气", &labels(&["天气", "邮件"])).unwrap();
552 assert_eq!(r[0].label, "天气");
553 }
554
555 #[test]
556 fn labels_that_cannot_be_told_apart_are_rejected() {
557 for bad in [
558 labels(&[]),
559 labels(&["--", "email"]),
560 labels(&["out_of_scope", "Out-of-scope"]),
561 ] {
562 assert!(
563 matches!(
564 score_reply("email", &bad),
565 Err(InferenceError::InvalidClassifyLabels(_))
566 ),
567 "{bad:?}"
568 );
569 }
570 }
571
572 #[test]
573 fn system_one_request_keeps_label_order_and_is_valid_json() {
574 let body = system_one_request_body(
575 "jev-1.13.0",
576 "move money to savings",
577 &labels(&["transfer", "out_of_scope", "balance"]),
578 )
579 .unwrap();
580 let positions: Vec<usize> = ["\"transfer\":", "\"out_of_scope\":", "\"balance\":"]
581 .iter()
582 .map(|k| body.find(k).unwrap())
583 .collect();
584 assert!(positions.windows(2).all(|w| w[0] < w[1]), "{body}");
585 let parsed: serde_json::Value = serde_json::from_str(&body).unwrap();
586 assert_eq!(parsed["model"], "jev-1.13.0");
587 assert_eq!(parsed["state"]["text"], "move money to savings");
588 assert_eq!(
589 parsed["questions"]["label"]["criteria"]["out_of_scope"],
590 "out of scope"
591 );
592 }
593
594 #[test]
595 fn system_one_results_are_probabilities_with_the_choice_first() {
596 let response = serde_json::json!({ "answers": { "label": {
597 "type": "choice", "choice": "balance",
598 "probabilities": { "transfer": 0.2, "balance": 0.7, "timer": 0.1 }
599 }}});
600 let r = system_one_results(&response, &labels(&["transfer", "balance", "timer"])).unwrap();
601 assert_eq!(r[0].label, "balance");
602 assert!((r[0].score - 0.7).abs() < 1e-9);
603 let foreign = serde_json::json!({ "answers": { "label": { "choice": "weather" }}});
604 assert!(system_one_results(&foreign, &labels(&["transfer", "balance"])).is_err());
605 let partial = serde_json::json!({ "answers": { "label": {
606 "choice": "balance", "probabilities": { "balance": 0.9 }
607 }}});
608 assert!(system_one_results(&partial, &labels(&["transfer", "balance"])).is_err());
609 }
610
611 #[test]
612 fn spellings_cover_the_capitalized_answer() {
613 assert_eq!(spellings("email"), ["email", "Email"]);
614 assert_eq!(spellings("Email"), ["Email"]);
615 assert!((log_sum_exp(&[0.5f64.ln(), 0.25f64.ln()]) - 0.75f64.ln()).abs() < 1e-12);
616 assert_eq!(log_sum_exp(&[f64::NEG_INFINITY]), f64::NEG_INFINITY);
617 }
618
619 #[test]
620 fn scores_are_normalized_match_strengths() {
621 let r = score_reply(
622 "book hotel",
623 &labels(&["book_hotel", "hotel_reviews", "timer"]),
624 )
625 .unwrap();
626 assert_eq!(r[0].label, "book_hotel");
627 let sum: f64 = r.iter().map(|x| x.score).sum();
628 assert!((sum - 1.0).abs() < 1e-9);
629 assert_eq!(r.last().unwrap().score, 0.0);
630 }
631}