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:
| Variant | Trigger | Display message ({} / to_string()) |
|---|---|---|
EmptyInput(String) | An array, vector, or dataset was empty where data was required | input is empty: <what> |
DimensionMismatch { expected, found } | 2 scalar counts disagreed | dimension 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 one | non-finite value (NaN or infinity) encountered in <where> |
InvalidParameter { name, reason } | A user-supplied hyperparameter was out of range | invalid 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 fit | model `<name>` has not been fitted; call `fit` before this operation |
NotConverged(String) | An iterative algorithm never met its convergence criterion | failed to converge: <msg> |
Computation { context, source } | A numerical breakdown, a violated invariant, or a wrapped foreign error | computation failed: <context> |
NeuralNetwork(NnError) | A neural-network-specific failure | forwarded transparently from NnError |
Tree(TreeError) | A decision-tree-specific failure | forwarded transparently from TreeError |
Io(IoError) | A filesystem or (de)serialization failure | forwarded 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)EmptyModelNotBuilt(&'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:
NotClassificationTreeCorruptStructure(&'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 failuresSerialization(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 layerUnsupportedModelFormat(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
metricsandmathmodules panic on an error instead of returningResult, which keeps those modules lightweight. - Whether an
ndarrayoperation outside RustyML returnsResultor panics isndarray’s decision, and RustyML has no say in it.