Keyboard shortcuts

Press or to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

1.6. Error Handling

1.6.1. 1 error type

RustyML has exactly 1 error type, rustyml::error::Error. Every fallible operation in the crate returns Result<T, rustyml::error::Error>, which also has an alias:

pub type RustymlResult<T> = std::result::Result<T, Error>;

Because Error is not in the prelude, you import the error machinery separately, with use rustyml::error::Error; and friends. These are Error’s variants:

VariantTriggerDisplay message ({} / to_string())
EmptyInput(String)An array, vector, or dataset was empty where data was requiredinput is empty: <what>
DimensionMismatch { expected, found }2 scalar counts disagreeddimension mismatch: expected <e>, found <f>
ShapeMismatch { expected, found }2 tensor shapes disagreed (a gradient vs. the activation it flows into)shape mismatch: expected [..], found [..]
NonFinite(String)A value in the data was NaN / inf, or a computation produced onenon-finite value (NaN or infinity) encountered in <where>
InvalidParameter { name, reason }A user-supplied hyperparameter was out of rangeinvalid parameter `<name>`: <reason>
InvalidInput(String)A validation failure with no more specific variant (bad rank, too few samples)invalid input: <msg>
NotFitted(&'static str)A method needing a trained model was called before fitmodel `<name>` has not been fitted; call `fit` before this operation
NotConverged(String)An iterative algorithm never met its convergence criterionfailed to converge: <msg>
Computation { context, source }A numerical breakdown, a violated invariant, or a wrapped foreign errorcomputation failed: <context>
NeuralNetwork(NnError)A neural-network-specific failureforwarded transparently from NnError
Tree(TreeError)A decision-tree-specific failureforwarded transparently from TreeError
Io(IoError)A filesystem or (de)serialization failureforwarded transparently from IoError

Note that DimensionMismatch compares scalar counts, such as a feature count or a vector length. ShapeMismatch is about 2 whole tensor shapes disagreeing, which shows up mostly in the neural-network code.

Error is annotated #[non_exhaustive], which means a match over it must carry a wildcard _ => (or Err(e) =>) arm.

1.6.2. The domain sub-errors

3 of Error’s variants each wrap a smaller enum. Concerns that only apply to neural networks (layer state, weight shapes, compilation) stay in their own enum. Concerns that only apply to trees (classification versus regression) stay in a separate enum of their own.

NnError (at rustyml::neural_network::NnError) contains:

  • ForwardPassNotRun(&'static str)
  • WeightShape { name, expected, found }
  • NotCompiled(&'static str)
  • EmptyModel
  • NotBuilt(&'static str)

NotBuilt is the newest of the 5. A layer allocates every array it owns in UnaryLayer::build, so a layer that never built holds no kernel and no bias. UnaryLayer::forward takes &self and cannot build, so it reports this variant. UnaryLayer::forward_mut takes &mut self, builds the layer from the tensor, and never raises it. A model that comes from SequentialBuilder::build is always built, and it never raises it.

Code example:

use rustyml::neural_network::Shape;
use rustyml::neural_network::sequential::SequentialBuilder;
use rustyml::neural_network::layers::Dense;
use rustyml::neural_network::layers::activation::ReLU;
use rustyml::neural_network::NnError;
use rustyml::error::Error;
use ndarray::Array;

fn main() {
    let mut model = SequentialBuilder::new()
        .add(Dense::new(2, ReLU::new()).unwrap())
        .build(&Shape::known(&[3, 4]))
        .unwrap();

    let x = Array::ones((3, 4)).into_dyn();
    let y = Array::ones((3, 2)).into_dyn();

    // compile() was never called, so no optimizer or loss is configured yet
    match model.fit(&x, &y, 1) {
        Ok(_) => unreachable!("training should not have started"),
        Err(Error::NeuralNetwork(NnError::NotCompiled(missing))) => {
            println!("compile the model first: `{missing}` is not specified");
        }
        Err(e) => println!("unexpected: {e}"),
    }
}

TreeError (at rustyml::machine_learning::TreeError) has these 2 variants:

  • NotClassificationTree
  • CorruptStructure(&'static str)

Code example:

use rustyml::machine_learning::{Algorithm, DecisionTree, TreeError};
use rustyml::error::Error;
use ndarray::array;

fn main() {
    // A regression tree (is_classifier = false) has no per-class probabilities
    let tree = DecisionTree::new(Algorithm::CART, false).unwrap();
    let x = array![[1.0, 2.0]];

    match tree.predict_proba(&x) {
        Err(Error::Tree(TreeError::NotClassificationTree)) => {
            println!("predict_proba is classification-only");
        }
        other => println!("unexpected: {other:?}"),
    }
}

IoError (at rustyml::error::IoError) has 4 variants:

  • Std(std::io::Error) for filesystem failures
  • Serialization(postcard::Error) for the binary format (RustyML serializes with postcard)
  • ModelStructureMismatch(String) for when a loaded neural-network file does not match the target architecture. Causes include a different number of layers, a different layer type at some position, or a weight whose shape does not fit the target layer
  • UnsupportedModelFormat(String) for when the file is not a RustyML model file at all, or its on-disk format version is not the one this build writes

Code example:

use rustyml::machine_learning::LinearRegression;
use rustyml::error::{Error, IoError};

fn main() {
    match LinearRegression::load_from_path("model_that_does_not_exist.bin") {
        Ok(_) => unreachable!("the file should not exist"),
        Err(Error::Io(IoError::Std(io_err))) => {
            // io_err is the underlying std::io::Error (kind NotFound here).
            println!("filesystem error: {io_err}");
        }
        Err(Error::Io(IoError::Serialization(e))) => {
            println!("the file exists but is not a valid model: {e}");
        }
        Err(e) => println!("unexpected: {e}"),
    }
}

For the serialization format and versioning, see 7.2. Model Persistence in Depth.

1.6.3. Matching on specific variants

The everyday failure is calling predict before fit, which returns Error::NotFitted carrying its own name as a &'static str:

use rustyml::machine_learning::LinearRegression;
use rustyml::error::Error;
use ndarray::array;

fn main() {
    // Constructed, but never fitted
    let model = LinearRegression::new(true);
    let x = array![[1.0, 2.0], [3.0, 4.0]];

    match model.predict(&x) {
        Ok(preds) => println!("{preds:?}"),
        Err(Error::NotFitted(name)) => {
            println!("`{name}` was not fitted; call fit() first");
        }
        Err(Error::DimensionMismatch { expected, found }) => {
            println!("wrong feature count: model wants {expected}, got {found}");
        }
        // `Error` is `#[non_exhaustive]`, so the wildcard arm is mandatory
        Err(e) => println!("other error: {e}"),
    }
}

The DimensionMismatch arm is there to show the pattern. This particular call actually triggers NotFitted. But if you feed a fitted model a matrix with the wrong number of columns, you take the second arm. expected is set to the feature count seen at fit time, and found is set to the one you passed to predict.

1.6.4. Propagating with ?

The whole crate uses a single error type, so you can return failures anywhere in a pipeline as Error with nothing beyond Result and ?:

use rustyml::machine_learning::{LinearRegression, RegularizationType};
use rustyml::error::RustymlResult;
use ndarray::{array, Array1, Array2};

fn train_and_predict(x: &Array2<f64>, y: &Array1<f64>) -> RustymlResult<Array1<f64>> {
    // Every ? below lifts a rustyml::error::Error out of a fallible call
    let mut model = LinearRegression::new(true)
        .with_regularization(RegularizationType::L2(0.01))?;       // maybe InvalidParameter
    model.fit(x, y)?;                                              // maybe EmptyInput / DimensionMismatch / NonFinite
    let preds = model.predict(x)?;                                 // maybe NotFitted / DimensionMismatch
    Ok(preds)
}

fn main() {
    let x = array![[1.0], [2.0], [3.0]];
    let y = Array1::from_vec(vec![2.0, 4.0, 6.0]);

    match train_and_predict(&x, &y) {
        Ok(preds) => println!("got {} predictions", preds.len()),
        Err(e) => eprintln!("pipeline failed: {e}"),
    }
}

When you do need to report a foreign error (from the standard library or another crate), reach for the Context extension trait. It lets you fold the foreign error into this scheme while keeping its cause chain, and you must import it into scope. RustyML implements it for any Result<T, E> whose E is Send + Sync + 'static and implements std::error::Error, so it composes with ?. context takes the message eagerly. with_context takes a closure that runs only on the error path. Whenever building the message allocates (anything with format!), prefer the closure form, so the success path never runs it:

use rustyml::error::{Context, Error, RustymlResult};

fn parse_threshold(raw: &str) -> RustymlResult<f64> {
    // A std ParseFloatError, wrapped together with our context as Error::Computation,
    // with its source() chain preserved for downcasting later.
    let value: f64 = raw
        .parse()
        .with_context(|| format!("parsing threshold from {raw:?}"))?;
    Ok(value)
}

fn main() {
    match parse_threshold("not-a-number") {
        Ok(v) => println!("threshold = {v}"),
        Err(Error::Computation { context, source }) => {
            println!("{context}");
            if let Some(cause) = source {
                println!("  caused by: {cause}");
            }
        }
        Err(e) => println!("unexpected: {e}"),
    }
}

The foreign error becomes the source of an Error::Computation, reachable through the standard std::error::Error::source() chain. It downcasts back to its original concrete type without losing any information.

1.6.5. Eager validation

RustyML’s error-handling design is that anything taking a hyperparameter validates it eagerly and returns Result, rather than panicking on illegal input.

use rustyml::machine_learning::LinearRegression;
use rustyml::machine_learning::linear_model::LeastSquaresSolver;
use rustyml::error::Error;

fn main() {
    // learning_rate must be positive and finite
    // 0.0 returns an error
    match LinearRegression::new(true).with_solver(LeastSquaresSolver::GradientDescent {
        learning_rate: 0.0,
        max_iter: 1000,
        tol: 1e-6,
    }) {
        Ok(_) => unreachable!("a zero learning rate must not be accepted"),
        Err(Error::InvalidParameter { name, reason }) => {
            // bad parameter `learning_rate`: must be positive and finite, got 0
            println!("bad parameter `{name}`: {reason}");
        }
        Err(e) => println!("unexpected: {e}"),
    }
}

A few places still panic outright:

  • The functions in the metrics and math modules panic on an error instead of returning Result, which keeps those modules lightweight.
  • Whether an ndarray operation outside RustyML returns Result or panics is ndarray’s decision, and RustyML has no say in it.