///|
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,
  )
}