use super::model::LRModel;
use super::serialization::{LRFormat, LRModelLoader};
#[derive(Debug, Clone)]
pub struct LRClassifier {
model: LRModel,
}
impl LRClassifier {
pub fn new() -> Self {
Self {
model: LRModel::new(),
}
}
pub fn from_model(model: LRModel) -> Self {
Self { model }
}
pub fn load(path: String) -> Result<Self, String> {
let loader = LRModelLoader::new();
let model = loader.load(path, LRFormat::Auto)?;
Ok(Self { model })
}
pub fn model(&self) -> &LRModel {
&self.model
}
pub fn predict(&self, features: &[String]) -> String {
let feature_ids = self.model.features_to_ids_readonly(features);
let logits = self.model.compute_logits(&feature_ids);
if logits.is_empty() {
return String::new();
}
let (best_class_id, _) = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((0, &0.0));
self.model
.id_to_class(best_class_id as u32)
.unwrap_or("")
.to_string()
}
pub fn predict_with_prob(&self, features: &[String]) -> (String, f64) {
let feature_ids = self.model.features_to_ids_readonly(features);
let logits = self.model.compute_logits(&feature_ids);
if logits.is_empty() {
return (String::new(), 0.0);
}
let probs = LRModel::softmax(&logits);
let (best_class_id, &best_prob) = probs
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((0, &0.0));
let class_label = self
.model
.id_to_class(best_class_id as u32)
.unwrap_or("")
.to_string();
(class_label, best_prob)
}
pub fn predict_proba(&self, features: &[String]) -> Vec<(String, f64)> {
let feature_ids = self.model.features_to_ids_readonly(features);
let logits = self.model.compute_logits(&feature_ids);
if logits.is_empty() {
return Vec::new();
}
let probs = LRModel::softmax(&logits);
let mut result: Vec<(String, f64)> = probs
.iter()
.enumerate()
.map(|(class_id, &prob)| {
let label = self
.model
.id_to_class(class_id as u32)
.unwrap_or("")
.to_string();
(label, prob)
})
.collect();
result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
result
}
pub fn predict_top_k(&self, features: &[String], k: usize) -> Vec<(String, f64)> {
let mut proba = self.predict_proba(features);
proba.truncate(k);
proba
}
pub fn num_classes(&self) -> usize {
self.model.num_classes
}
pub fn classes(&self) -> Vec<String> {
self.model.get_classes()
}
}
impl Default for LRClassifier {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_model() -> LRModel {
let mut model = LRModel::with_classes(vec![
"positive".to_string(),
"negative".to_string(),
"neutral".to_string(),
]);
let f1 = model.features.get_or_insert("word=good");
let f2 = model.features.get_or_insert("word=bad");
let f3 = model.features.get_or_insert("word=okay");
model.num_features = model.features.len();
model.set_weight(f1, 0, 2.0); model.set_weight(f1, 1, -1.0); model.set_weight(f2, 1, 2.0); model.set_weight(f2, 0, -1.0); model.set_weight(f3, 2, 2.0);
model
}
#[test]
fn test_classifier_new() {
let classifier = LRClassifier::new();
assert_eq!(classifier.num_classes(), 0);
}
#[test]
fn test_classifier_from_model() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
assert_eq!(classifier.num_classes(), 3);
}
#[test]
fn test_predict() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let pred = classifier.predict(&["word=good".to_string()]);
assert_eq!(pred, "positive");
let pred = classifier.predict(&["word=bad".to_string()]);
assert_eq!(pred, "negative");
let pred = classifier.predict(&["word=okay".to_string()]);
assert_eq!(pred, "neutral");
}
#[test]
fn test_predict_with_prob() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let (label, prob) = classifier.predict_with_prob(&["word=good".to_string()]);
assert_eq!(label, "positive");
assert!(prob > 0.5, "Probability should be > 0.5, got {}", prob);
}
#[test]
fn test_predict_proba() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let proba = classifier.predict_proba(&["word=good".to_string()]);
assert_eq!(proba.len(), 3);
for i in 0..proba.len() - 1 {
assert!(proba[i].1 >= proba[i + 1].1);
}
let sum: f64 = proba.iter().map(|(_, p)| p).sum();
assert!((sum - 1.0).abs() < 1e-6);
assert_eq!(proba[0].0, "positive");
}
#[test]
fn test_predict_top_k() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let top2 = classifier.predict_top_k(&["word=good".to_string()], 2);
assert_eq!(top2.len(), 2);
assert_eq!(top2[0].0, "positive");
}
#[test]
fn test_predict_unknown_features() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let pred = classifier.predict(&["word=unknown".to_string()]);
assert!(!pred.is_empty());
}
#[test]
fn test_get_classes() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let classes = classifier.classes();
assert_eq!(classes.len(), 3);
assert!(classes.contains(&"positive".to_string()));
assert!(classes.contains(&"negative".to_string()));
assert!(classes.contains(&"neutral".to_string()));
}
#[test]
fn test_empty_features() {
let model = create_test_model();
let classifier = LRClassifier::from_model(model);
let pred = classifier.predict(&[]);
assert!(!pred.is_empty());
let proba = classifier.predict_proba(&[]);
assert_eq!(proba.len(), 3);
}
}