///|
pub fn MulticlassModel::classes(self : MulticlassModel) -> Array[Int] {
self.ordered_classes.copy()
}
///|
pub fn MulticlassModel::model_count(self : MulticlassModel) -> Int {
self.class_models.length()
}
///|
pub fn MulticlassModel::feature_count(self : MulticlassModel) -> Int {
self.input_columns
}
///|
pub fn MulticlassModel::binary_models(
self : MulticlassModel,
) -> Array[BinaryModel] {
self.class_models.copy()
}
///|
pub fn MulticlassModel::decision_values(
self : MulticlassModel,
row : Array[Double],
) -> Result[Array[Double], SvmError] {
if row.length() != self.input_columns {
return Err(PredictionDimensionMismatch(self.input_columns, row.length()))
}
let scores : Array[Double] = []
for model in self.class_models {
match model.decision_value(row) {
Err(error) => return Err(error)
Ok(score) => scores.push(score)
}
}
Ok(scores)
}
///|
/// Chooses the largest one-vs-rest score; exact ties retain sorted class order.
pub fn MulticlassModel::predict(
self : MulticlassModel,
row : Array[Double],
) -> Result[Int, SvmError] {
let scores = match self.decision_values(row) {
Err(error) => return Err(error)
Ok(value) => value
}
let mut best = 0
for index = 1; index < scores.length(); index = index + 1 {
if scores[index] > scores[best] {
best = index
}
}
Ok(self.ordered_classes[best])
}
///|
pub fn MulticlassModel::predict_batch(
self : MulticlassModel,
rows : Array[Array[Double]],
) -> Result[Array[Int], SvmError] {
let predictions : Array[Int] = []
for row in rows {
match self.predict(row) {
Err(error) => return Err(error)
Ok(label) => predictions.push(label)
}
}
Ok(predictions)
}
///|
fn one_vs_rest_dataset(
data : Dataset,
positive_class : Int,
) -> Result[Dataset, SvmError] {
let binary_labels = Array::make(data.row_count(), -1)
for index, label in data.labels() {
if label == positive_class {
binary_labels[index] = 1
}
}
dataset(data.features(), binary_labels, "one-vs-rest class \{positive_class}")
}
///|
/// Fits one deterministic weighted binary classifier for each sorted class.
pub fn train_multiclass_weighted(
data : Dataset,
sample_weights : Array[Double],
class_weights : Array[ClassWeight],
config : BinaryConfig,
) -> Result[MulticlassModel, SvmError] {
let classes = data.classes()
if classes.length() < 2 {
return Err(SingleClassDataset)
}
if sample_weights.length() != data.row_count() {
return Err(WeightLengthMismatch(data.row_count(), sample_weights.length()))
}
for index, value in sample_weights {
if !finite_double(value) || value <= 0.0 {
return Err(InvalidWeight(index, value))
}
}
match validate_class_weights(classes, class_weights) {
Err(error) => return Err(error)
Ok(_) => ()
}
let original_labels = data.labels()
let effective_weights = Array::make(data.row_count(), 0.0)
for index, value in sample_weights {
let effective = value *
class_multiplier(original_labels[index], class_weights)
if !finite_double(effective) || effective <= 0.0 {
return Err(ZeroEffectiveWeight)
}
effective_weights[index] = effective
}
let models : Array[BinaryModel] = []
for positive_class in classes {
let binary_data = match one_vs_rest_dataset(data, positive_class) {
Err(error) => return Err(error)
Ok(value) => value
}
let model = match
train_binary_weighted(binary_data, effective_weights, [], config) {
Err(error) => return Err(error)
Ok(value) => value
}
models.push(model)
}
Ok({
ordered_classes: classes,
class_models: models,
input_columns: data.feature_count(),
})
}
///|
/// Fits an unweighted deterministic one-vs-rest multiclass classifier.
pub fn train_multiclass(
data : Dataset,
config : BinaryConfig,
) -> Result[MulticlassModel, SvmError] {
train_multiclass_weighted(
data,
Array::make(data.row_count(), 1.0),
[],
config,
)
}