Skip to main content

car_inference/tasks/
classify.rs

1//! Classification — score text against candidate labels using prompt-based inference.
2
3use 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/// A classification request.
12#[derive(Debug, Clone, Serialize, Deserialize)]
13pub struct ClassifyRequest {
14    /// The text to classify.
15    pub text: String,
16    /// Candidate labels to score against.
17    pub labels: Vec<String>,
18    /// Optional model override.
19    pub model: Option<String>,
20    /// Trusted lifetime context, supplied by the caller rather than model input.
21    #[serde(skip)]
22    pub work_context: Option<car_auth::context::CredentialContext>,
23}
24
25/// A classification result with label and confidence score.
26#[derive(Debug, Clone, Serialize, Deserialize)]
27pub struct ClassifyResult {
28    pub label: String,
29    pub score: f64,
30}
31
32/// The result of [`crate::InferenceEngine::option_probabilities`].
33#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct OptionProbabilities {
35    /// Per option, in the order given, renormalized to sum to 1.
36    pub probabilities: Vec<f64>,
37    /// Probability the options received before renormalizing. Low means the
38    /// model wanted to answer something else.
39    pub mass: f64,
40    /// `"first_token"` or `"sequence"`: which scoring ran (see the method docs).
41    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/// "email" and "Email": the spellings a chat model is likely to start its
70/// answer with.
71#[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/// Score `options` as the start of the assistant's answer to `formatted` (an
89/// already-rendered chat prompt). See `InferenceEngine::option_probabilities`.
90#[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    // Raw tokenization for the spliced answer: no special tokens, so a
98    // tokenizer that adds a BOS cannot make every option start the same.
99    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
171/// Lowercase, with every run of non-alphanumeric characters as one space, so
172/// `out_of_scope`, `Out-of-scope.` and `out of scope` compare equal.
173fn 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
181/// Words that carry no label meaning on their own. A refusal such as "I
182/// can't classify this" must not partially match "are you a bot" on "a".
183const 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
188/// Scripts written without spaces between words, where whole-word matching
189/// would never find a label inside a reply.
190fn unspaced_script(text: &str) -> bool {
191    text.chars().any(|c| {
192        matches!(c as u32,
193            0x3040..=0x30FF   // Hiragana, Katakana
194            | 0x3400..=0x4DBF // CJK Extension A
195            | 0x4E00..=0x9FFF // CJK Unified Ideographs
196            | 0x0E00..=0x0E7F // Thai
197            | 0x0E80..=0x0EFF // Lao
198            | 0x1780..=0x17FF // Khmer
199            | 0x1000..=0x109F // Myanmar
200        )
201    })
202}
203
204/// Reject labels that cannot be told apart before asking anything.
205pub 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
228/// Score each label against a generative model's reply to the classify prompt.
229///
230/// Scores are match strengths normalized to sum to 1, not probabilities:
231/// 1.0 when the reply is the label (or the label's number, since the prompt
232/// numbers them), 0.8 when the reply contains it as whole words (as a
233/// substring for scripts written without spaces), otherwise half the share of
234/// the label's content words the reply uses, counted only when that share is
235/// at least one half.
236///
237/// A reply that names no label, or names several equally without naming one
238/// exactly, is [`InferenceError::ClassifyNoAnswer`] — never a list whose top
239/// entry is whichever label happened to be offered first.
240pub 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
320/// The System One request for one classify call, serialized by hand so the
321/// criteria keep the caller's label order (a `serde_json::Map` may sort keys).
322/// The criterion text is the label made readable — the same information a
323/// generative model's prompt gives it.
324pub 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
350/// Every label with the probability System One gave it, best first. A
351/// response missing the answer, or naming a label that was not offered, is
352/// an error rather than a guess.
353pub 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    // Every label must carry the model's own probability; a missing one is not
376    // filled in, since the docs promise these scores are the model's.
377    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    // The chosen label stays first even if rounding ties it with another.
400    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))))]
408/// Classify text against candidate labels.
409///
410/// Uses a prompt-based approach: asks the model to pick the best label,
411/// then parses the response with [`score_reply`].
412pub 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, // greedy for classification
438            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        // An exact answer is never a tie.
534        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}