use serde::{Deserialize, Serialize};
#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
use crate::backend::CandleBackend;
#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
use crate::tasks::generate;
use crate::InferenceError;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClassifyRequest {
pub text: String,
pub labels: Vec<String>,
pub model: Option<String>,
#[serde(skip)]
pub work_context: Option<car_auth::context::CredentialContext>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClassifyResult {
pub label: String,
pub score: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OptionProbabilities {
pub probabilities: Vec<f64>,
pub mass: f64,
pub method: String,
}
#[cfg(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx)))]
fn log_prob(logits: &[f32], token: u32) -> Result<f64, InferenceError> {
let at = logits.get(token as usize).ok_or_else(|| {
InferenceError::InferenceFailed(format!(
"token {token} is outside the model's {}-entry vocabulary",
logits.len()
))
})?;
let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
let sum: f64 = logits.iter().map(|&l| ((l as f64) - max).exp()).sum();
Ok((*at as f64) - max - sum.ln())
}
#[cfg(any(
test,
all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))
))]
fn log_sum_exp(values: &[f64]) -> f64 {
let max = values.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
if max == f64::NEG_INFINITY {
return max;
}
max + values.iter().map(|v| (v - max).exp()).sum::<f64>().ln()
}
#[cfg(any(
test,
all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))
))]
fn spellings(option: &str) -> Vec<String> {
let mut chars = option.chars();
let capitalized = match chars.next() {
Some(c) => c.to_uppercase().collect::<String>() + chars.as_str(),
None => String::new(),
};
let mut out = vec![option.to_string()];
if capitalized != option {
out.push(capitalized);
}
out
}
#[cfg(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx)))]
pub fn score_options(
backend: &mut crate::backend::SwiftLmBackend,
formatted: &str,
options: &[String],
) -> Result<OptionProbabilities, InferenceError> {
let prompt = backend.encode(formatted)?;
let variants: Vec<Vec<Vec<u32>>> = options
.iter()
.map(|o| {
spellings(o)
.iter()
.map(|v| backend.tokenize_raw(v))
.collect::<Result<Vec<_>, _>>()
})
.collect::<Result<_, _>>()?;
if variants.iter().flatten().any(|t| t.is_empty()) {
return Err(InferenceError::InvalidClassifyLabels(
"an option encodes to no tokens".into(),
));
}
let firsts: Vec<std::collections::HashSet<u32>> = variants
.iter()
.map(|v| v.iter().map(|t| t[0]).collect())
.collect();
let collide = firsts
.iter()
.enumerate()
.any(|(i, a)| firsts.iter().skip(i + 1).any(|b| !a.is_disjoint(b)));
let (scores, method) = if !collide {
backend.clear_kv_cache();
let logits = backend.forward(&prompt, 0)?;
let scores = firsts
.iter()
.map(|ids| {
let lps = ids
.iter()
.map(|&id| log_prob(&logits, id))
.collect::<Result<Vec<_>, _>>()?;
Ok(log_sum_exp(&lps))
})
.collect::<Result<Vec<_>, InferenceError>>()?;
(scores, "first_token")
} else {
let end = backend.token_id("<|im_end|>");
let mut scores = Vec::with_capacity(options.len());
for option in &variants {
let mut per_spelling = Vec::with_capacity(option.len());
for tokens in option {
let sequence: Vec<u32> = tokens.iter().copied().chain(end).collect();
backend.clear_kv_cache();
let mut logits = backend.forward(&prompt, 0)?;
let mut total = 0.0;
for (i, &token) in sequence.iter().enumerate() {
total += log_prob(&logits, token)?;
if i + 1 < sequence.len() {
logits = backend.forward(&[token], prompt.len() + i)?;
}
}
per_spelling.push(total);
}
scores.push(log_sum_exp(&per_spelling));
}
(scores, "sequence")
};
let mass: f64 = scores.iter().map(|s| s.exp()).sum();
if mass.is_nan() || mass <= 0.0 {
return Err(InferenceError::InferenceFailed(
"the model gave none of the options any probability".into(),
));
}
let top = log_sum_exp(&scores);
Ok(OptionProbabilities {
probabilities: scores.iter().map(|s| (s - top).exp()).collect(),
mass,
method: method.into(),
})
}
fn normalize(text: &str) -> String {
text.to_lowercase()
.split(|c: char| !c.is_alphanumeric())
.filter(|w| !w.is_empty())
.collect::<Vec<_>>()
.join(" ")
}
const STOPWORDS: &[&str] = &[
"a", "an", "the", "i", "is", "are", "am", "be", "of", "to", "in", "on", "it", "this", "that",
"and", "or", "for", "with", "my", "me", "you", "your", "can", "t", "do", "not",
];
fn unspaced_script(text: &str) -> bool {
text.chars().any(|c| {
matches!(c as u32,
0x3040..=0x30FF | 0x3400..=0x4DBF | 0x4E00..=0x9FFF | 0x0E00..=0x0E7F | 0x0E80..=0x0EFF | 0x1780..=0x17FF | 0x1000..=0x109F )
})
}
pub fn validate_labels(labels: &[String]) -> Result<(), InferenceError> {
if labels.is_empty() {
return Err(InferenceError::InvalidClassifyLabels(
"no labels given".into(),
));
}
let mut seen = std::collections::HashMap::new();
for label in labels {
let norm = normalize(label);
if norm.is_empty() {
return Err(InferenceError::InvalidClassifyLabels(format!(
"{label:?} has no letters or digits"
)));
}
if let Some(other) = seen.insert(norm, label) {
return Err(InferenceError::InvalidClassifyLabels(format!(
"{other:?} and {label:?} differ only in case or separators"
)));
}
}
Ok(())
}
pub fn score_reply(reply: &str, labels: &[String]) -> Result<Vec<ClassifyResult>, InferenceError> {
validate_labels(labels)?;
let no_answer = |reason: &str| InferenceError::ClassifyNoAnswer {
reply: reply.trim().chars().take(120).collect(),
reason: reason.into(),
};
let reply_norm = normalize(reply);
if let Ok(n) = reply_norm.parse::<usize>() {
if (1..=labels.len()).contains(&n) {
return Ok(labels
.iter()
.enumerate()
.map(|(i, label)| ClassifyResult {
label: label.clone(),
score: if i + 1 == n { 1.0 } else { 0.0 },
})
.collect::<Vec<_>>())
.map(|mut results: Vec<ClassifyResult>| {
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results
});
}
}
let reply_words: std::collections::HashSet<&str> = reply_norm.split(' ').collect();
let padded = format!(" {reply_norm} ");
let mut results: Vec<ClassifyResult> = labels
.iter()
.map(|label| {
let label_norm = normalize(label);
let contained = padded.contains(&format!(" {label_norm} "))
|| (unspaced_script(&label_norm) && reply_norm.contains(&label_norm));
let score = if reply_norm == label_norm {
1.0
} else if contained {
0.8
} else {
let content: Vec<&str> = label_norm
.split(' ')
.filter(|w| !STOPWORDS.contains(w))
.collect();
let hits = content.iter().filter(|w| reply_words.contains(*w)).count();
let share = if content.is_empty() {
0.0
} else {
hits as f64 / content.len() as f64
};
if share >= 0.5 {
0.5 * share
} else {
0.0
}
};
ClassifyResult {
label: label.clone(),
score,
}
})
.collect();
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let total: f64 = results.iter().map(|r| r.score).sum();
if total <= 0.0 {
return Err(no_answer("names none of the labels"));
}
if results.len() > 1 && results[0].score < 1.0 && results[0].score == results[1].score {
return Err(no_answer("names several labels equally"));
}
for r in &mut results {
r.score /= total;
}
Ok(results)
}
pub fn system_one_request_body(
model: &str,
text: &str,
labels: &[String],
) -> Result<String, InferenceError> {
validate_labels(labels)?;
let criteria = labels
.iter()
.map(|label| {
format!(
"{}:{}",
serde_json::Value::from(label.as_str()),
serde_json::Value::from(normalize(label))
)
})
.collect::<Vec<_>>()
.join(",");
Ok(format!(
"{{\"model\":{},\"state\":{},\"questions\":{{\"label\":{{\"type\":\"choice\",\
\"instructions\":{},\"criteria\":{{{criteria}}}}}}}}}",
serde_json::Value::from(model),
serde_json::json!({ "text": text }),
serde_json::Value::from("Classify `text` into one of the labels."),
))
}
pub fn system_one_results(
response: &serde_json::Value,
labels: &[String],
) -> Result<Vec<ClassifyResult>, InferenceError> {
let answer = response.pointer("/answers/label").ok_or_else(|| {
InferenceError::InferenceFailed(format!("System One response has no answer: {response}"))
})?;
let choice = answer
.get("choice")
.and_then(|c| c.as_str())
.ok_or_else(|| InferenceError::InferenceFailed("System One answer has no choice".into()))?;
if !labels.iter().any(|l| l == choice) {
return Err(InferenceError::InferenceFailed(format!(
"System One chose {choice:?}, which is not an offered label"
)));
}
let probabilities = answer
.get("probabilities")
.and_then(|p| p.as_object())
.ok_or_else(|| {
InferenceError::InferenceFailed("System One answer has no probabilities".into())
})?;
let mut results = labels
.iter()
.map(|label| {
probabilities
.get(label)
.and_then(|p| p.as_f64())
.map(|score| ClassifyResult {
label: label.clone(),
score,
})
.ok_or_else(|| {
InferenceError::InferenceFailed(format!(
"System One gave no probability for {label:?}"
))
})
})
.collect::<Result<Vec<_>, _>>()?;
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
if let Some(i) = results.iter().position(|r| r.label == choice) {
let chosen = results.remove(i);
results.insert(0, chosen);
}
Ok(results)
}
#[cfg(not(all(target_os = "macos", target_arch = "aarch64", not(car_skip_mlx))))]
pub async fn classify(
backend: &mut CandleBackend,
req: ClassifyRequest,
) -> Result<Vec<ClassifyResult>, InferenceError> {
let labels_str = req
.labels
.iter()
.enumerate()
.map(|(i, l)| format!("{}. {}", i + 1, l))
.collect::<Vec<_>>()
.join("\n");
let prompt = format!(
"Classify the following text into one of these categories:\n\
{labels_str}\n\n\
Text: {}\n\n\
Respond with ONLY the category name, nothing else.",
req.text
);
let gen_req = generate::GenerateRequest {
work_context: req.work_context.clone(),
prompt,
model: req.model.clone(),
params: generate::GenerateParams {
temperature: 0.0, max_tokens: 32,
..Default::default()
},
context: None,
context_stable_prefix: None,
tools: None,
images: None,
messages: None,
cache_control: false,
response_format: None,
intent: None,
client_ref: None,
expected_row_digest: None,
expected_catalog_revision: None,
caller: None,
};
let (response, _ttft_ms, _prompt_tokens, _completion_tokens) =
generate::generate(backend, gen_req).await?;
score_reply(&response, &req.labels)
.map_err(|e| InferenceError::InferenceFailed(format!("classify: {e}")))
}
#[cfg(test)]
mod score_tests {
use super::*;
#[test]
fn classification_context_is_trusted_and_not_serialized() {
let context = car_auth::context::CredentialContext {
api_base: "https://authority.example".into(),
account_id: "original-account".into(),
organization_id: Some("parslee".into()),
};
let request = ClassifyRequest {
text: "a task".into(),
labels: vec!["work".into()],
model: None,
work_context: Some(context.clone()),
};
let mut serialized = serde_json::to_value(&request).unwrap();
assert!(serialized.get("work_context").is_none());
serialized["work_context"] = serde_json::to_value(context).unwrap();
let decoded: ClassifyRequest = serde_json::from_value(serialized).unwrap();
assert!(decoded.work_context.is_none());
}
fn labels(names: &[&str]) -> Vec<String> {
names.iter().map(|s| s.to_string()).collect()
}
fn no_answer(r: Result<Vec<ClassifyResult>, InferenceError>) -> String {
match r {
Err(InferenceError::ClassifyNoAnswer { reason, .. }) => reason,
other => panic!("expected ClassifyNoAnswer, got {other:?}"),
}
}
#[test]
fn separators_and_case_do_not_hide_a_label() {
let r = score_reply("Out of scope.", &labels(&["transfer", "out_of_scope"])).unwrap();
assert_eq!(r[0].label, "out_of_scope");
assert!(r[0].score > 0.99);
}
#[test]
fn a_label_inside_a_reply_is_matched_on_whole_words() {
let r = score_reply(
"I think it is transfer money",
&labels(&["transfer", "yes"]),
)
.unwrap();
assert_eq!(r[0].label, "transfer");
assert_eq!(
no_answer(score_reply("yesterday", &labels(&["yes", "no"]))),
"names none of the labels"
);
}
#[test]
fn a_refusal_does_not_match_on_small_words() {
let r = score_reply(
"I can't classify this",
&labels(&["are you a bot", "book hotel"]),
);
assert_eq!(no_answer(r), "names none of the labels");
}
#[test]
fn several_labels_named_equally_is_no_answer() {
let r = score_reply(
"transfer or balance",
&labels(&["transfer", "balance", "timer"]),
);
assert_eq!(no_answer(r), "names several labels equally");
let r = score_reply("hotel", &labels(&["book hotel", "hotel", "hotel reviews"])).unwrap();
assert_eq!(r[0].label, "hotel");
}
#[test]
fn a_numbered_reply_names_that_label() {
let r = score_reply("2.", &labels(&["email", "calendar", "search"])).unwrap();
assert_eq!(r[0].label, "calendar");
assert_eq!(r[0].score, 1.0);
assert_eq!(
no_answer(score_reply("7", &labels(&["email", "calendar"]))),
"names none of the labels"
);
}
#[test]
fn an_unspaced_script_label_is_found_inside_the_reply() {
let r = score_reply("今天天气", &labels(&["天气", "邮件"])).unwrap();
assert_eq!(r[0].label, "天气");
}
#[test]
fn labels_that_cannot_be_told_apart_are_rejected() {
for bad in [
labels(&[]),
labels(&["--", "email"]),
labels(&["out_of_scope", "Out-of-scope"]),
] {
assert!(
matches!(
score_reply("email", &bad),
Err(InferenceError::InvalidClassifyLabels(_))
),
"{bad:?}"
);
}
}
#[test]
fn system_one_request_keeps_label_order_and_is_valid_json() {
let body = system_one_request_body(
"jev-1.13.0",
"move money to savings",
&labels(&["transfer", "out_of_scope", "balance"]),
)
.unwrap();
let positions: Vec<usize> = ["\"transfer\":", "\"out_of_scope\":", "\"balance\":"]
.iter()
.map(|k| body.find(k).unwrap())
.collect();
assert!(positions.windows(2).all(|w| w[0] < w[1]), "{body}");
let parsed: serde_json::Value = serde_json::from_str(&body).unwrap();
assert_eq!(parsed["model"], "jev-1.13.0");
assert_eq!(parsed["state"]["text"], "move money to savings");
assert_eq!(
parsed["questions"]["label"]["criteria"]["out_of_scope"],
"out of scope"
);
}
#[test]
fn system_one_results_are_probabilities_with_the_choice_first() {
let response = serde_json::json!({ "answers": { "label": {
"type": "choice", "choice": "balance",
"probabilities": { "transfer": 0.2, "balance": 0.7, "timer": 0.1 }
}}});
let r = system_one_results(&response, &labels(&["transfer", "balance", "timer"])).unwrap();
assert_eq!(r[0].label, "balance");
assert!((r[0].score - 0.7).abs() < 1e-9);
let foreign = serde_json::json!({ "answers": { "label": { "choice": "weather" }}});
assert!(system_one_results(&foreign, &labels(&["transfer", "balance"])).is_err());
let partial = serde_json::json!({ "answers": { "label": {
"choice": "balance", "probabilities": { "balance": 0.9 }
}}});
assert!(system_one_results(&partial, &labels(&["transfer", "balance"])).is_err());
}
#[test]
fn spellings_cover_the_capitalized_answer() {
assert_eq!(spellings("email"), ["email", "Email"]);
assert_eq!(spellings("Email"), ["Email"]);
assert!((log_sum_exp(&[0.5f64.ln(), 0.25f64.ln()]) - 0.75f64.ln()).abs() < 1e-12);
assert_eq!(log_sum_exp(&[f64::NEG_INFINITY]), f64::NEG_INFINITY);
}
#[test]
fn scores_are_normalized_match_strengths() {
let r = score_reply(
"book hotel",
&labels(&["book_hotel", "hotel_reviews", "timer"]),
)
.unwrap();
assert_eq!(r[0].label, "book_hotel");
let sum: f64 = r.iter().map(|x| x.score).sum();
assert!((sum - 1.0).abs() < 1e-9);
assert_eq!(r.last().unwrap().score, 0.0);
}
}