use crate::error::{Error, Result};
use crate::frame::{Dataset, Frame};
use crate::traits::{Balancer, Estimator, Model, ParamValue, Predictor, Transformer};
#[derive(Default, Clone)]
pub struct Pipeline {
steps: Vec<(String, Box<dyn Transformer>)>,
balancer: Option<Box<dyn Balancer>>,
estimator: Option<(String, Box<dyn Model>)>,
fitted: bool,
}
impl Pipeline {
pub fn new() -> Self {
Pipeline {
steps: Vec::new(),
balancer: None,
estimator: None,
fitted: false,
}
}
pub fn step(mut self, name: impl Into<String>, t: impl Transformer + 'static) -> Self {
self.steps.push((name.into(), Box::new(t)));
self
}
pub fn balance(mut self, b: impl Balancer + 'static) -> Self {
self.balancer = Some(Box::new(b));
self
}
pub fn estimator(mut self, name: impl Into<String>, e: impl Model + 'static) -> Self {
self.estimator = Some((name.into(), Box::new(e)));
self
}
pub fn step_names(&self) -> Vec<&str> {
self.steps
.iter()
.map(|(n, _)| n.as_str())
.chain(self.estimator.iter().map(|(n, _)| n.as_str()))
.collect()
}
pub fn set_param(&mut self, path: &str, value: ParamValue) -> Result<()> {
let (step, rest) = path
.split_once("__")
.ok_or_else(|| Error::Param(format!("'{path}' is not a 'step__param' path")))?;
for (name, t) in &mut self.steps {
if name == step {
return t.set_param(rest, value);
}
}
if let Some((name, e)) = &mut self.estimator {
if name == step {
return e.set_param(rest, value);
}
}
Err(Error::Param(format!("no step named '{step}' in pipeline")))
}
fn require_estimator(&self) -> Result<&(String, Box<dyn Model>)> {
self.estimator
.as_ref()
.ok_or_else(|| Error::Pipeline("pipeline has no estimator".into()))
}
fn forward(&self, frame: &Frame) -> Result<Frame> {
let mut current = frame.clone();
for (_, t) in &self.steps {
current = t.transform(¤t)?;
}
Ok(current)
}
}
impl Estimator for Pipeline {
fn name(&self) -> &'static str {
"Pipeline"
}
fn fit(&mut self, dataset: &Dataset) -> Result<()> {
self.require_estimator()?;
let mut current = dataset.features().clone();
for (_, t) in &mut self.steps {
current = t.fit_transform(¤t)?;
}
let transformed = match &self.balancer {
Some(b) => {
let (bx, by) = b.fit_resample(¤t, dataset.target())?;
Dataset::new(bx, by)?
}
None => dataset.with_features(current),
};
let (_, est) = self.estimator.as_mut().ok_or_else(|| {
Error::Pipeline("pipeline has no estimator; add one before fit".into())
})?;
est.fit(&transformed)?;
self.fitted = true;
Ok(())
}
fn set_param(&mut self, name: &str, value: ParamValue) -> Result<()> {
Pipeline::set_param(self, name, value)
}
#[cfg(feature = "onnx")]
fn to_onnx_proto(&self) -> Result<onnx_export_rs::proto::ModelProto> {
crate::onnx::ExportOnnx::to_onnx(self)
}
}
impl Predictor for Pipeline {
fn predict(&self, frame: &Frame) -> Result<Vec<f64>> {
if !self.fitted {
return Err(Error::NotFitted("Pipeline::predict".into()));
}
let transformed = self.forward(frame)?;
let (_, est) = self.require_estimator()?;
est.predict(&transformed)
}
}
#[cfg(feature = "onnx")]
impl crate::onnx::ExportOnnx for Pipeline {
fn to_onnx(&self) -> Result<onnx_export_rs::proto::ModelProto> {
let (_, est) = self.require_estimator()?;
let mut proto = est.to_onnx_proto()?;
let prefixes: Vec<crate::onnx::Prefix> = self
.steps
.iter()
.map(|(name, t)| {
t.onnx_prefix().ok_or_else(|| {
Error::Backend(format!(
"pipeline step '{name}' ({}) is not ONNX-exportable",
t.name()
))
})
})
.collect::<Result<_>>()?;
if !prefixes.is_empty() {
crate::onnx::prepend_prefixes(&mut proto, &prefixes)?;
}
Ok(proto)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transform::StandardScaler;
#[derive(Default, Clone)]
struct MeanBaseline {
mean: f64,
}
impl Estimator for MeanBaseline {
fn name(&self) -> &'static str {
"MeanBaseline"
}
fn fit(&mut self, d: &Dataset) -> Result<()> {
self.mean = d.target().iter().sum::<f64>() / d.target().len() as f64;
Ok(())
}
}
impl Predictor for MeanBaseline {
fn predict(&self, frame: &Frame) -> Result<Vec<f64>> {
Ok(vec![self.mean; frame.nrows()])
}
}
#[test]
fn fits_and_predicts_through_a_transform() {
let x = Frame::from_rows(
vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]],
vec!["x".into()],
)
.unwrap();
let ds = Dataset::new(x.clone(), vec![10.0, 10.0, 20.0, 20.0]).unwrap();
let mut pipe = Pipeline::new()
.step("scale", StandardScaler::new())
.estimator("mean", MeanBaseline::default());
pipe.fit(&ds).unwrap();
let preds = pipe.predict(&x).unwrap();
assert_eq!(preds, vec![15.0, 15.0, 15.0, 15.0]);
assert_eq!(pipe.step_names(), vec!["scale", "mean"]);
}
#[test]
fn routes_params_by_path() {
let mut pipe = Pipeline::new()
.step("scale", StandardScaler::new())
.estimator("mean", MeanBaseline::default());
assert!(pipe
.set_param("scale__with_mean", ParamValue::Bool(false))
.is_ok());
assert!(pipe.set_param("nope__x", ParamValue::Int(1)).is_err());
assert!(pipe.set_param("scale__bogus", ParamValue::Int(1)).is_err());
assert!(pipe.set_param("noseparator", ParamValue::Int(1)).is_err());
}
#[test]
fn predict_before_fit_errors() {
let x = Frame::from_rows(vec![vec![1.0]], vec!["x".into()]).unwrap();
let pipe = Pipeline::new().estimator("mean", MeanBaseline::default());
assert!(pipe.predict(&x).is_err());
}
}