onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use crate::canonical::{AggregationMode, ForestStructure, GradientBoostedEnsemble};
use crate::exporters::tree_ensemble::export_tree_ensemble_with_options;
use crate::proto::ModelProto;
use crate::{Error, Result};

/// Exports gradient boosting as a summed ONNX-ML tree ensemble with base scores.
///
/// # Errors
///
/// Returns an error for an empty ensemble, inconsistent target/base-score
/// counts, non-finite learning rate, or malformed tree leaves.
pub fn export_gradient_boosting(model: &GradientBoostedEnsemble) -> Result<ModelProto> {
    if model.trees.is_empty()
        || model.n_targets == 0
        || model.base_values.len() != model.n_targets
        || !model.learning_rate.is_finite()
    {
        return Err(Error::InvalidModel(
            "gradient-boosted ensemble shape mismatch".into(),
        ));
    }
    let mut trees = model.trees.clone();
    for tree in &mut trees {
        for node in &mut tree.nodes {
            for value in &mut node.leaf_values {
                *value *= model.learning_rate as f32;
            }
        }
    }
    let forest = ForestStructure {
        trees,
        aggregation: AggregationMode::Sum,
        n_targets: model.n_targets,
    };
    let base_values: Vec<f32> = model
        .base_values
        .iter()
        .map(|&value| value as f32)
        .collect();
    export_tree_ensemble_with_options(
        &forest,
        model.task,
        &base_values,
        model.post_transform.as_onnx(),
    )
}