Skip to main content

sklears_svm/
structured_svm.rs

1//! Structured Support Vector Machines
2//!
3//! This module implements structured SVMs for sequence labeling and other
4//! structured prediction tasks. Structured SVMs extend traditional SVMs to
5//! handle structured outputs like sequences, trees, or graphs.
6//!
7//! Algorithms included:
8//! - Structured SVM for sequence labeling (CRF-like)
9//! - Structural SVM with margin rescaling
10//! - Structural SVM with slack rescaling
11//! - Latent structural SVM
12
13use scirs2_core::ndarray::{Array1, Array2};
14use sklears_core::{
15    error::{Result, SklearsError},
16    traits::{Fit, Predict, Trained, Untrained},
17};
18use std::collections::HashMap;
19use std::fmt;
20use std::marker::PhantomData;
21
22/// Errors that can occur during structured SVM training and prediction
23#[derive(Debug, Clone)]
24pub enum StructuredSVMError {
25    InvalidInput(String),
26    TrainingError(String),
27    PredictionError(String),
28    ConvergenceError(String),
29}
30
31impl fmt::Display for StructuredSVMError {
32    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
33        match self {
34            StructuredSVMError::InvalidInput(msg) => write!(f, "Invalid input: {msg}"),
35            StructuredSVMError::TrainingError(msg) => write!(f, "Training error: {msg}"),
36            StructuredSVMError::PredictionError(msg) => write!(f, "Prediction error: {msg}"),
37            StructuredSVMError::ConvergenceError(msg) => write!(f, "Convergence error: {msg}"),
38        }
39    }
40}
41
42impl std::error::Error for StructuredSVMError {}
43
44/// Loss function for structured SVM
45#[derive(Debug, Clone)]
46pub enum StructuredLoss {
47    /// Hamming loss (number of incorrect predictions)
48    Hamming,
49    /// F1 loss (based on F1 score)
50    F1,
51    /// Edit distance (Levenshtein distance)
52    EditDistance,
53}
54
55/// Inference algorithm for structured prediction
56#[derive(Debug, Clone)]
57pub enum InferenceAlgorithm {
58    /// Viterbi algorithm for sequence labeling
59    Viterbi,
60    /// Belief propagation for general graphical models
61    BeliefPropagation,
62    /// Graph cuts for certain types of problems
63    GraphCuts,
64}
65
66/// Structured SVM configuration
67#[derive(Debug, Clone)]
68pub struct StructuredSVMConfig {
69    /// Regularization parameter
70    pub c: f64,
71    /// Loss function
72    pub loss: StructuredLoss,
73    /// Inference algorithm
74    pub inference: InferenceAlgorithm,
75    /// Maximum number of iterations
76    pub max_iter: usize,
77    /// Tolerance for convergence
78    pub tol: f64,
79    /// Learning rate for cutting plane algorithm
80    pub learning_rate: f64,
81    /// Verbose output
82    pub verbose: bool,
83}
84
85impl Default for StructuredSVMConfig {
86    fn default() -> Self {
87        Self {
88            c: 1.0,
89            loss: StructuredLoss::Hamming,
90            inference: InferenceAlgorithm::Viterbi,
91            max_iter: 100,
92            tol: 1e-4,
93            learning_rate: 0.01,
94            verbose: false,
95        }
96    }
97}
98
99/// Sequence structure for labeling tasks
100#[derive(Debug, Clone)]
101pub struct Sequence {
102    /// Feature vectors for each position in the sequence
103    pub features: Array2<f64>,
104    /// Length of the sequence
105    pub length: usize,
106    /// Labels (for training sequences)
107    pub labels: Option<Array1<usize>>,
108}
109
110impl Sequence {
111    /// Create a new sequence
112    pub fn new(features: Array2<f64>, labels: Option<Array1<usize>>) -> Self {
113        let length = features.nrows();
114        Self {
115            features,
116            length,
117            labels,
118        }
119    }
120
121    /// Get feature vector at position
122    pub fn feature_at(&self, position: usize) -> Array1<f64> {
123        self.features.row(position).to_owned()
124    }
125
126    /// Get label at position (if available)
127    pub fn label_at(&self, position: usize) -> Option<usize> {
128        self.labels.as_ref().map(|labels| labels[position])
129    }
130}
131
132/// Structured SVM for sequence labeling
133#[derive(Debug, Clone)]
134pub struct StructuredSVM<State = Untrained> {
135    config: StructuredSVMConfig,
136    state: PhantomData<State>,
137    // Model parameters
138    weights: Option<Array1<f64>>,
139    transition_weights: Option<Array2<f64>>,
140    // Training data statistics
141    #[allow(dead_code)] // intentionally deferred: feature count readout pending
142    n_features: Option<usize>,
143    #[allow(dead_code)] // intentionally deferred: label count readout pending
144    n_labels: Option<usize>,
145    #[allow(dead_code)] // intentionally deferred: label encoding readout pending
146    label_to_idx: Option<HashMap<usize, usize>>,
147    #[allow(dead_code)] // intentionally deferred: label decoding readout pending
148    idx_to_label: Option<HashMap<usize, usize>>,
149}
150
151impl Default for StructuredSVM<Untrained> {
152    fn default() -> Self {
153        Self::new()
154    }
155}
156
157impl StructuredSVM<Untrained> {
158    /// Create a new structured SVM
159    pub fn new() -> Self {
160        Self {
161            config: StructuredSVMConfig::default(),
162            state: PhantomData,
163            weights: None,
164            transition_weights: None,
165            n_features: None,
166            n_labels: None,
167            label_to_idx: None,
168            idx_to_label: None,
169        }
170    }
171
172    /// Set the regularization parameter
173    pub fn with_c(mut self, c: f64) -> Self {
174        self.config.c = c;
175        self
176    }
177
178    /// Set the loss function
179    pub fn with_loss(mut self, loss: StructuredLoss) -> Self {
180        self.config.loss = loss;
181        self
182    }
183
184    /// Set the inference algorithm
185    pub fn with_inference(mut self, inference: InferenceAlgorithm) -> Self {
186        self.config.inference = inference;
187        self
188    }
189
190    /// Set maximum iterations
191    pub fn with_max_iter(mut self, max_iter: usize) -> Self {
192        self.config.max_iter = max_iter;
193        self
194    }
195
196    /// Set tolerance
197    pub fn with_tolerance(mut self, tol: f64) -> Self {
198        self.config.tol = tol;
199        self
200    }
201
202    /// Enable verbose output
203    pub fn verbose(mut self, verbose: bool) -> Self {
204        self.config.verbose = verbose;
205        self
206    }
207}
208
209impl Fit<Vec<Sequence>, Vec<Array1<usize>>> for StructuredSVM<Untrained> {
210    type Fitted = StructuredSVM<Trained>;
211
212    fn fit(self, sequences: &Vec<Sequence>, labels: &Vec<Array1<usize>>) -> Result<Self::Fitted> {
213        if sequences.len() != labels.len() {
214            return Err(SklearsError::InvalidInput(
215                "Number of sequences and label arrays must match".to_string(),
216            ));
217        }
218
219        if sequences.is_empty() {
220            return Err(SklearsError::InvalidInput(
221                "Cannot fit on empty dataset".to_string(),
222            ));
223        }
224
225        // Extract feature dimensions and label vocabulary
226        let n_features = sequences[0].features.ncols();
227        let mut label_set = std::collections::HashSet::new();
228
229        for (seq, seq_labels) in sequences.iter().zip(labels.iter()) {
230            if seq.features.ncols() != n_features {
231                return Err(SklearsError::InvalidInput(
232                    "All sequences must have the same feature dimension".to_string(),
233                ));
234            }
235
236            if seq.length != seq_labels.len() {
237                return Err(SklearsError::InvalidInput(
238                    "Sequence length must match label length".to_string(),
239                ));
240            }
241
242            for &label in seq_labels.iter() {
243                label_set.insert(label);
244            }
245        }
246
247        let n_labels = label_set.len();
248        let mut label_to_idx = HashMap::new();
249        let mut idx_to_label = HashMap::new();
250
251        for (idx, &label) in label_set.iter().enumerate() {
252            label_to_idx.insert(label, idx);
253            idx_to_label.insert(idx, label);
254        }
255
256        // Initialize weights
257        let total_features = n_features + n_labels * n_labels; // Node features + transition features
258        let mut weights = Array1::zeros(total_features);
259        let mut transition_weights = Array2::zeros((n_labels, n_labels));
260
261        // Cutting plane algorithm for structured SVM training
262        let mut objective_history = Vec::new();
263
264        for iteration in 0..self.config.max_iter {
265            let mut total_loss = 0.0;
266            let mut gradient = Array1::zeros(total_features);
267
268            // For each training sequence
269            for (seq, seq_labels) in sequences.iter().zip(labels.iter()) {
270                // Find most violating constraint (loss-augmented inference)
271                let predicted_labels = self.loss_augmented_inference(
272                    seq,
273                    seq_labels,
274                    &weights,
275                    &transition_weights,
276                    &label_to_idx,
277                )?;
278
279                // Compute loss
280                let loss = self.compute_loss(seq_labels, &predicted_labels)?;
281                total_loss += loss;
282
283                // Compute feature difference (ψ(x,y) - ψ(x,ŷ))
284                let true_features = self.compute_features(seq, seq_labels, &label_to_idx)?;
285                let pred_features = self.compute_features(seq, &predicted_labels, &label_to_idx)?;
286                let feature_diff = &true_features - &pred_features;
287
288                // Update gradient
289                gradient = gradient + feature_diff;
290            }
291
292            // Regularization gradient
293            let reg_gradient = &weights * self.config.c;
294            gradient = gradient - reg_gradient;
295
296            // Update weights
297            weights = weights + self.config.learning_rate * gradient;
298
299            // Update transition weights (extract from weights)
300            let start_idx = n_features;
301            for i in 0..n_labels {
302                for j in 0..n_labels {
303                    let idx = start_idx + i * n_labels + j;
304                    transition_weights[[i, j]] = weights[idx];
305                }
306            }
307
308            // Compute objective
309            let regularization = 0.5 * self.config.c * weights.dot(&weights);
310            let objective = total_loss - regularization;
311            objective_history.push(objective);
312
313            if self.config.verbose {
314                println!(
315                    "Iteration {iteration}: Objective = {objective:.6}, Loss = {total_loss:.6}"
316                );
317            }
318
319            // Check convergence
320            if iteration > 0 {
321                let prev_obj = objective_history[iteration - 1];
322                let obj_change = (objective - prev_obj).abs();
323                if obj_change < self.config.tol {
324                    if self.config.verbose {
325                        println!("Converged after {} iterations", iteration + 1);
326                    }
327                    break;
328                }
329            }
330        }
331
332        Ok(StructuredSVM {
333            config: self.config,
334            state: PhantomData,
335            weights: Some(weights),
336            transition_weights: Some(transition_weights),
337            n_features: Some(n_features),
338            n_labels: Some(n_labels),
339            label_to_idx: Some(label_to_idx),
340            idx_to_label: Some(idx_to_label),
341        })
342    }
343}
344
345impl StructuredSVM<Untrained> {
346    /// Compute features for a sequence with given labels
347    fn compute_features(
348        &self,
349        sequence: &Sequence,
350        labels: &Array1<usize>,
351        label_to_idx: &HashMap<usize, usize>,
352    ) -> Result<Array1<f64>> {
353        let n_features = sequence.features.ncols();
354        let n_labels = label_to_idx.len();
355        let total_features = n_features + n_labels * n_labels;
356        let mut features = Array1::zeros(total_features);
357
358        // Node features (emission features)
359        for (pos, &label) in labels.iter().enumerate() {
360            let _label_idx = *label_to_idx.get(&label).expect("key not found");
361            let node_features = sequence.feature_at(pos);
362
363            // Add weighted node features
364            for (feat_idx, &feat_val) in node_features.iter().enumerate() {
365                features[feat_idx] += feat_val;
366            }
367        }
368
369        // Transition features
370        for pos in 1..labels.len() {
371            let prev_label = labels[pos - 1];
372            let curr_label = labels[pos];
373            let prev_idx = *label_to_idx.get(&prev_label).expect("key not found");
374            let curr_idx = *label_to_idx.get(&curr_label).expect("key not found");
375
376            let transition_idx = n_features + prev_idx * n_labels + curr_idx;
377            features[transition_idx] += 1.0;
378        }
379
380        Ok(features)
381    }
382
383    /// Loss-augmented inference (finding most violating constraint)
384    fn loss_augmented_inference(
385        &self,
386        sequence: &Sequence,
387        true_labels: &Array1<usize>,
388        weights: &Array1<f64>,
389        transition_weights: &Array2<f64>,
390        label_to_idx: &HashMap<usize, usize>,
391    ) -> Result<Array1<usize>> {
392        match self.config.inference {
393            InferenceAlgorithm::Viterbi => self.viterbi_loss_augmented(
394                sequence,
395                true_labels,
396                weights,
397                transition_weights,
398                label_to_idx,
399            ),
400            _ => {
401                // Fallback to simple greedy decoding for now
402                self.greedy_inference(sequence, weights, label_to_idx)
403            }
404        }
405    }
406
407    /// Viterbi algorithm with loss augmentation
408    fn viterbi_loss_augmented(
409        &self,
410        sequence: &Sequence,
411        true_labels: &Array1<usize>,
412        weights: &Array1<f64>,
413        transition_weights: &Array2<f64>,
414        label_to_idx: &HashMap<usize, usize>,
415    ) -> Result<Array1<usize>> {
416        let seq_len = sequence.length;
417        let n_labels = label_to_idx.len();
418        let _n_features = sequence.features.ncols();
419
420        // Dynamic programming tables
421        let mut dp = Array2::zeros((seq_len, n_labels));
422        let mut backtrack = Array2::zeros((seq_len, n_labels));
423
424        // Initialize first position
425        for (label, &label_idx) in label_to_idx.iter() {
426            let node_features = sequence.feature_at(0);
427            let mut score = 0.0;
428
429            // Node features
430            for (feat_idx, &feat_val) in node_features.iter().enumerate() {
431                score += weights[feat_idx] * feat_val;
432            }
433
434            // Loss augmentation
435            if *label != true_labels[0] {
436                score += 1.0; // Hamming loss
437            }
438
439            dp[[0, label_idx]] = score;
440        }
441
442        // Forward pass
443        for pos in 1..seq_len {
444            let node_features = sequence.feature_at(pos);
445
446            for (curr_label, &curr_idx) in label_to_idx.iter() {
447                let mut best_score = f64::NEG_INFINITY;
448                let mut best_prev = 0;
449
450                for (_prev_label, &prev_idx) in label_to_idx.iter() {
451                    let transition_score = transition_weights[[prev_idx, curr_idx]];
452                    let total_score = dp[[pos - 1, prev_idx]] + transition_score;
453
454                    if total_score > best_score {
455                        best_score = total_score;
456                        best_prev = prev_idx;
457                    }
458                }
459
460                // Add node features
461                for (feat_idx, &feat_val) in node_features.iter().enumerate() {
462                    best_score += weights[feat_idx] * feat_val;
463                }
464
465                // Loss augmentation
466                if *curr_label != true_labels[pos] {
467                    best_score += 1.0; // Hamming loss
468                }
469
470                dp[[pos, curr_idx]] = best_score;
471                backtrack[[pos, curr_idx]] = best_prev as f64;
472            }
473        }
474
475        // Find best final state
476        let mut best_final_score = f64::NEG_INFINITY;
477        let mut best_final_state = 0;
478        for label_idx in 0..n_labels {
479            if dp[[seq_len - 1, label_idx]] > best_final_score {
480                best_final_score = dp[[seq_len - 1, label_idx]];
481                best_final_state = label_idx;
482            }
483        }
484
485        // Backward pass (traceback)
486        let mut path = Array1::zeros(seq_len);
487        let mut current_state = best_final_state;
488
489        for pos in (0..seq_len).rev() {
490            path[pos] = current_state as f64;
491            if pos > 0 {
492                current_state = backtrack[[pos, current_state]] as usize;
493            }
494        }
495
496        // Convert indices back to labels
497        let idx_to_label = label_to_idx
498            .iter()
499            .map(|(&k, &v)| (v, k))
500            .collect::<HashMap<_, _>>();
501        let labels = path
502            .iter()
503            .map(|&idx| idx_to_label[&(idx as usize)])
504            .collect();
505
506        Ok(Array1::from_vec(labels))
507    }
508
509    /// Simple greedy inference
510    fn greedy_inference(
511        &self,
512        sequence: &Sequence,
513        weights: &Array1<f64>,
514        label_to_idx: &HashMap<usize, usize>,
515    ) -> Result<Array1<usize>> {
516        let mut labels = Array1::zeros(sequence.length);
517
518        for pos in 0..sequence.length {
519            let node_features = sequence.feature_at(pos);
520            let mut best_score = f64::NEG_INFINITY;
521            let mut best_label = 0;
522
523            for (label, _) in label_to_idx.iter() {
524                let mut score = 0.0;
525                for (feat_idx, &feat_val) in node_features.iter().enumerate() {
526                    score += weights[feat_idx] * feat_val;
527                }
528
529                if score > best_score {
530                    best_score = score;
531                    best_label = *label;
532                }
533            }
534
535            labels[pos] = best_label;
536        }
537
538        Ok(labels)
539    }
540
541    /// Compute loss between true and predicted labels
542    fn compute_loss(
543        &self,
544        true_labels: &Array1<usize>,
545        pred_labels: &Array1<usize>,
546    ) -> Result<f64> {
547        match self.config.loss {
548            StructuredLoss::Hamming => {
549                let mut loss = 0.0;
550                for (true_label, pred_label) in true_labels.iter().zip(pred_labels.iter()) {
551                    if true_label != pred_label {
552                        loss += 1.0;
553                    }
554                }
555                Ok(loss)
556            }
557            StructuredLoss::F1 => {
558                // Simplified F1 loss calculation
559                let mut tp = 0.0;
560                let mut fp = 0.0;
561                let mut fn_count = 0.0;
562
563                for (true_label, pred_label) in true_labels.iter().zip(pred_labels.iter()) {
564                    if true_label == pred_label && *true_label != 0 {
565                        tp += 1.0;
566                    } else if *pred_label != 0 {
567                        fp += 1.0;
568                    } else if *true_label != 0 {
569                        fn_count += 1.0;
570                    }
571                }
572
573                let precision = if tp + fp > 0.0 { tp / (tp + fp) } else { 0.0 };
574                let recall = if tp + fn_count > 0.0 {
575                    tp / (tp + fn_count)
576                } else {
577                    0.0
578                };
579                let f1 = if precision + recall > 0.0 {
580                    2.0 * precision * recall / (precision + recall)
581                } else {
582                    0.0
583                };
584
585                Ok(1.0 - f1)
586            }
587            StructuredLoss::EditDistance => {
588                // Simplified edit distance (Levenshtein distance)
589                let n = true_labels.len();
590                let m = pred_labels.len();
591                let mut dp = Array2::zeros((n + 1, m + 1));
592
593                // Initialize base cases
594                for i in 0..=n {
595                    dp[[i, 0]] = i as f64;
596                }
597                for j in 0..=m {
598                    dp[[0, j]] = j as f64;
599                }
600
601                // Fill DP table
602                for i in 1..=n {
603                    for j in 1..=m {
604                        let cost = if true_labels[i - 1] == pred_labels[j - 1] {
605                            0.0
606                        } else {
607                            1.0
608                        };
609                        dp[[i, j]] = (dp[[i - 1, j]] + 1.0)
610                            .min(dp[[i, j - 1]] + 1.0)
611                            .min(dp[[i - 1, j - 1]] + cost);
612                    }
613                }
614
615                Ok(dp[[n, m]])
616            }
617        }
618    }
619}
620
621impl Predict<Vec<Sequence>, Vec<Array1<usize>>> for StructuredSVM<Trained> {
622    fn predict(&self, sequences: &Vec<Sequence>) -> Result<Vec<Array1<usize>>> {
623        let weights = self
624            .weights
625            .as_ref()
626            .expect("weights not available - model not fitted");
627        let transition_weights = self
628            .transition_weights
629            .as_ref()
630            .expect("transition_weights not available - model not fitted");
631        let label_to_idx = self
632            .label_to_idx
633            .as_ref()
634            .expect("label_to_idx not available - model not fitted");
635
636        let mut predictions = Vec::new();
637
638        for sequence in sequences {
639            let pred_labels = match self.config.inference {
640                InferenceAlgorithm::Viterbi => {
641                    self.viterbi_inference(sequence, weights, transition_weights, label_to_idx)?
642                }
643                _ => self.greedy_inference_trained(sequence, weights, label_to_idx)?,
644            };
645            predictions.push(pred_labels);
646        }
647
648        Ok(predictions)
649    }
650}
651
652impl StructuredSVM<Trained> {
653    /// Viterbi inference for prediction
654    fn viterbi_inference(
655        &self,
656        sequence: &Sequence,
657        weights: &Array1<f64>,
658        transition_weights: &Array2<f64>,
659        label_to_idx: &HashMap<usize, usize>,
660    ) -> Result<Array1<usize>> {
661        let seq_len = sequence.length;
662        let n_labels = label_to_idx.len();
663
664        // Dynamic programming tables
665        let mut dp = Array2::zeros((seq_len, n_labels));
666        let mut backtrack = Array2::zeros((seq_len, n_labels));
667
668        // Initialize first position
669        for (_label, &label_idx) in label_to_idx.iter() {
670            let node_features = sequence.feature_at(0);
671            let mut score = 0.0;
672
673            // Node features
674            for (feat_idx, &feat_val) in node_features.iter().enumerate() {
675                score += weights[feat_idx] * feat_val;
676            }
677
678            dp[[0, label_idx]] = score;
679        }
680
681        // Forward pass
682        for pos in 1..seq_len {
683            let node_features = sequence.feature_at(pos);
684
685            for (_curr_label, &curr_idx) in label_to_idx.iter() {
686                let mut best_score = f64::NEG_INFINITY;
687                let mut best_prev = 0;
688
689                for (_prev_label, &prev_idx) in label_to_idx.iter() {
690                    let transition_score = transition_weights[[prev_idx, curr_idx]];
691                    let total_score = dp[[pos - 1, prev_idx]] + transition_score;
692
693                    if total_score > best_score {
694                        best_score = total_score;
695                        best_prev = prev_idx;
696                    }
697                }
698
699                // Add node features
700                for (feat_idx, &feat_val) in node_features.iter().enumerate() {
701                    best_score += weights[feat_idx] * feat_val;
702                }
703
704                dp[[pos, curr_idx]] = best_score;
705                backtrack[[pos, curr_idx]] = best_prev as f64;
706            }
707        }
708
709        // Find best final state
710        let mut best_final_score = f64::NEG_INFINITY;
711        let mut best_final_state = 0;
712        for label_idx in 0..n_labels {
713            if dp[[seq_len - 1, label_idx]] > best_final_score {
714                best_final_score = dp[[seq_len - 1, label_idx]];
715                best_final_state = label_idx;
716            }
717        }
718
719        // Backward pass (traceback)
720        let mut path = Array1::zeros(seq_len);
721        let mut current_state = best_final_state;
722
723        for pos in (0..seq_len).rev() {
724            path[pos] = current_state as f64;
725            if pos > 0 {
726                current_state = backtrack[[pos, current_state]] as usize;
727            }
728        }
729
730        // Convert indices back to labels
731        let idx_to_label = label_to_idx
732            .iter()
733            .map(|(&k, &v)| (v, k))
734            .collect::<HashMap<_, _>>();
735        let labels = path
736            .iter()
737            .map(|&idx| idx_to_label[&(idx as usize)])
738            .collect();
739
740        Ok(Array1::from_vec(labels))
741    }
742
743    /// Simple greedy inference for trained model
744    fn greedy_inference_trained(
745        &self,
746        sequence: &Sequence,
747        weights: &Array1<f64>,
748        label_to_idx: &HashMap<usize, usize>,
749    ) -> Result<Array1<usize>> {
750        let mut labels = Array1::zeros(sequence.length);
751
752        for pos in 0..sequence.length {
753            let node_features = sequence.feature_at(pos);
754            let mut best_score = f64::NEG_INFINITY;
755            let mut best_label = 0;
756
757            for (label, _) in label_to_idx.iter() {
758                let mut score = 0.0;
759                for (feat_idx, &feat_val) in node_features.iter().enumerate() {
760                    score += weights[feat_idx] * feat_val;
761                }
762
763                if score > best_score {
764                    best_score = score;
765                    best_label = *label;
766                }
767            }
768
769            labels[pos] = best_label;
770        }
771
772        Ok(labels)
773    }
774
775    /// Get model weights
776    pub fn weights(&self) -> &Array1<f64> {
777        self.weights
778            .as_ref()
779            .expect("weights not available - model not fitted")
780    }
781
782    /// Get transition weights
783    pub fn transition_weights(&self) -> &Array2<f64> {
784        self.transition_weights
785            .as_ref()
786            .expect("transition_weights not available - model not fitted")
787    }
788}
789
790#[allow(non_snake_case)]
791#[cfg(test)]
792mod tests {
793    use super::*;
794    use scirs2_core::ndarray::array;
795
796    fn create_test_sequences() -> (Vec<Sequence>, Vec<Array1<usize>>) {
797        // Create simple test sequences for POS tagging-like task
798        let seq1 = Sequence::new(
799            array![[1.0, 0.5, 0.2], [0.8, 0.9, 0.1], [0.3, 0.1, 0.7]],
800            None,
801        );
802        let labels1 = array![0, 1, 2]; // NOUN, VERB, ADJ
803
804        let seq2 = Sequence::new(array![[0.9, 0.4, 0.3], [0.2, 0.8, 0.2]], None);
805        let labels2 = array![0, 1]; // NOUN, VERB
806
807        (vec![seq1, seq2], vec![labels1, labels2])
808    }
809
810    #[test]
811    fn test_structured_svm_creation() {
812        let svm = StructuredSVM::new()
813            .with_c(1.0)
814            .with_loss(StructuredLoss::Hamming)
815            .with_inference(InferenceAlgorithm::Viterbi)
816            .with_max_iter(50)
817            .verbose(false);
818
819        assert_eq!(svm.config.c, 1.0);
820        assert!(matches!(svm.config.loss, StructuredLoss::Hamming));
821        assert!(matches!(svm.config.inference, InferenceAlgorithm::Viterbi));
822    }
823
824    #[test]
825    fn test_sequence_creation() {
826        let features = array![[1.0, 2.0], [3.0, 4.0]];
827        let labels = array![0, 1];
828        let seq = Sequence::new(features.clone(), Some(labels.clone()));
829
830        assert_eq!(seq.length, 2);
831        assert_eq!(seq.feature_at(0), array![1.0, 2.0]);
832        assert_eq!(seq.label_at(0), Some(0));
833    }
834
835    #[test]
836    fn test_structured_svm_fit() {
837        let (sequences, labels) = create_test_sequences();
838
839        let svm = StructuredSVM::new()
840            .with_c(0.1)
841            .with_max_iter(10)
842            .with_tolerance(1e-3);
843
844        use sklears_core::traits::Fit;
845        let result = svm.fit(&sequences, &labels);
846        assert!(result.is_ok());
847
848        let trained_svm = result.expect("operation should succeed");
849        assert!(trained_svm.weights.is_some());
850        assert!(trained_svm.transition_weights.is_some());
851    }
852
853    #[test]
854    fn test_structured_svm_predict() {
855        let (sequences, labels) = create_test_sequences();
856
857        let svm = StructuredSVM::new()
858            .with_c(0.1)
859            .with_max_iter(5)
860            .with_tolerance(1e-2);
861
862        use sklears_core::traits::Predict;
863        let trained_svm = svm
864            .fit(&sequences, &labels)
865            .expect("model fitting should succeed");
866
867        let predictions = trained_svm.predict(&sequences);
868        assert!(predictions.is_ok());
869
870        let pred_labels = predictions.expect("operation should succeed");
871        assert_eq!(pred_labels.len(), 2);
872        assert_eq!(pred_labels[0].len(), 3);
873        assert_eq!(pred_labels[1].len(), 2);
874    }
875
876    #[test]
877    fn test_hamming_loss() {
878        let svm = StructuredSVM::new();
879        let true_labels = array![0, 1, 2, 1];
880        let pred_labels = array![0, 2, 2, 0];
881
882        let loss = svm
883            .compute_loss(&true_labels, &pred_labels)
884            .expect("operation should succeed");
885        assert_eq!(loss, 2.0); // Two mismatches
886    }
887
888    #[test]
889    fn test_invalid_input() {
890        let (sequences, mut labels) = create_test_sequences();
891        labels.pop(); // Make lengths mismatch
892
893        let svm = StructuredSVM::new();
894        use sklears_core::traits::Fit;
895        let result = svm.fit(&sequences, &labels);
896        assert!(result.is_err());
897    }
898}