use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use dataflow_rs::datavalue::{DType, OwnedDataTensor};
use serde::Serialize;
use super::manifest::Manifest;
use crate::config::ModelsConfig;
pub mod tract;
pub use self::tract::TractRuntime;
pub const NAMES: &[&str] = &["tract"];
pub const KNOWN_FORMATS: &[&str] = &["onnx"];
const TRACT_DEVICES: &[&str] = &["cpu", "metal", "cuda"];
const TRACT_FORMATS: &[&str] = &["onnx"];
pub fn devices_of(name: &str) -> Option<&'static [&'static str]> {
match name {
"tract" => Some(TRACT_DEVICES),
_ => None,
}
}
pub fn formats_of(name: &str) -> Option<&'static [&'static str]> {
match name {
"tract" => Some(TRACT_FORMATS),
_ => None,
}
}
pub fn intern(name: &str) -> Option<&'static str> {
NAMES.iter().copied().find(|n| *n == name)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoadError {
pub stage: &'static str,
pub message: String,
}
impl LoadError {
pub fn new(stage: &'static str, message: impl Into<String>) -> Self {
Self {
stage,
message: message.into(),
}
}
}
impl fmt::Display for LoadError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.stage, self.message)
}
}
impl std::error::Error for LoadError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RunError {
pub stage: &'static str,
pub message: String,
}
impl RunError {
pub fn new(stage: &'static str, message: impl Into<String>) -> Self {
Self {
stage,
message: message.into(),
}
}
}
impl fmt::Display for RunError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.stage, self.message)
}
}
impl std::error::Error for RunError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RuntimeSelection {
NoDefault { format: String },
UnknownRuntime { runtime: String, format: String },
NotEnabled { runtime: String, format: String },
DoesNotServe {
runtime: String,
format: String,
serves: &'static [&'static str],
},
}
impl fmt::Display for RuntimeSelection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NoDefault { format } => write!(
f,
"no default runtime for format '{format}' (models.default_runtime.{format} is \
not set)"
),
Self::UnknownRuntime { runtime, format } => write!(
f,
"runtime '{runtime}' for format '{format}' is unknown: this build lists {}",
NAMES.join(", ")
),
Self::NotEnabled { runtime, format } => write!(
f,
"runtime '{runtime}' for format '{format}' is not enabled on this node \
(models.runtimes.{runtime}.enabled)"
),
Self::DoesNotServe {
runtime,
format,
serves,
} => write!(
f,
"runtime '{runtime}' does not serve format '{format}'; it serves {}",
serves
.iter()
.map(|s| format!("'{s}'"))
.collect::<Vec<_>>()
.join(", ")
),
}
}
}
impl std::error::Error for RuntimeSelection {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoadBinding {
inputs: Vec<BoundInput>,
outputs: Vec<String>,
fingerprint: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct BoundInput {
pub name: String,
pub dtype: String,
pub shape: Vec<super::manifest::Dim>,
}
impl BoundInput {
pub fn dtype(&self) -> Option<DType> {
super::manifest::parse_dtype(&self.dtype).ok()
}
}
impl LoadBinding {
pub fn of(manifest: &Manifest) -> Self {
#[derive(Serialize)]
struct Rendered<'a> {
inputs: &'a [BoundInput],
outputs: &'a [String],
}
let inputs: Vec<BoundInput> = manifest
.inputs
.iter()
.map(|input| BoundInput {
name: input.name.clone(),
dtype: input.dtype.clone(),
shape: input.shape.clone(),
})
.collect();
let outputs: Vec<String> = manifest.output_names().map(str::to_string).collect();
let rendered = serde_json::to_vec(&Rendered {
inputs: &inputs,
outputs: &outputs,
})
.expect("a binding of strings and numbers serializes");
Self {
fingerprint: crate::crypto::sha256_digest(&rendered),
inputs,
outputs,
}
}
pub fn inputs(&self) -> &[BoundInput] {
&self.inputs
}
pub fn input_names(&self) -> impl Iterator<Item = &str> {
self.inputs.iter().map(|input| input.name.as_str())
}
pub fn output_names(&self) -> impl Iterator<Item = &str> {
self.outputs.iter().map(String::as_str)
}
pub fn fingerprint(&self) -> &str {
&self.fingerprint
}
}
pub trait ModelRuntime: Send + Sync {
fn name(&self) -> &'static str;
fn devices(&self) -> &'static [&'static str];
fn formats(&self) -> &'static [&'static str];
fn load(
&self,
bytes: &[u8],
binding: &LoadBinding,
device: &str,
) -> Result<Arc<dyn LoadedModel>, LoadError>;
}
pub trait LoadedModel: Send + Sync {
fn digest(&self) -> &str;
fn resident_bytes(&self) -> usize;
fn run(&self, inputs: Vec<OwnedDataTensor>) -> Result<Vec<OwnedDataTensor>, RunError>;
}
#[derive(Default)]
pub struct ModelRuntimes {
by_name: BTreeMap<&'static str, Arc<dyn ModelRuntime>>,
}
impl ModelRuntimes {
pub fn empty() -> Self {
Self::default()
}
pub fn builtin(config: &ModelsConfig) -> Self {
let mut registry = Self::empty();
if config
.runtimes
.get(self::tract::NAME)
.is_some_and(|runtime| runtime.enabled)
{
registry
.by_name
.insert(self::tract::NAME, Arc::new(TractRuntime));
}
registry
}
pub fn register(&mut self, runtime: Arc<dyn ModelRuntime>) -> Result<(), String> {
let name = runtime.name();
if !NAMES.contains(&name) {
return Err(format!(
"runtime '{name}' is not one this build lists ({}); add it to \
model::runtimes::NAMES before registering it",
NAMES.join(", ")
));
}
if runtime.formats().is_empty() {
return Err(format!(
"runtime '{name}' serves no format, so nothing could ever select it"
));
}
if self.by_name.contains_key(name) {
return Err(format!("runtime '{name}' is registered twice"));
}
self.by_name.insert(name, runtime);
Ok(())
}
pub fn get(&self, name: &str) -> Option<Arc<dyn ModelRuntime>> {
self.by_name.get(name).cloned()
}
pub fn default_for<'c>(
&self,
config: &'c ModelsConfig,
format: &str,
) -> Result<(Arc<dyn ModelRuntime>, &'c str), RuntimeSelection> {
let Some(name) = config.default_runtime_for(format) else {
return Err(RuntimeSelection::NoDefault {
format: format.to_string(),
});
};
self.for_format(config, name, format)
}
pub fn for_format<'c>(
&self,
config: &'c ModelsConfig,
name: &str,
format: &str,
) -> Result<(Arc<dyn ModelRuntime>, &'c str), RuntimeSelection> {
if !NAMES.contains(&name) {
return Err(RuntimeSelection::UnknownRuntime {
runtime: name.to_string(),
format: format.to_string(),
});
}
let enabled = config
.runtimes
.get(name)
.is_some_and(|runtime| runtime.enabled);
let (Some(runtime), Some(device)) =
(self.get(name).filter(|_| enabled), config.device_of(name))
else {
return Err(RuntimeSelection::NotEnabled {
runtime: name.to_string(),
format: format.to_string(),
});
};
let serves = runtime.formats();
if !serves.contains(&format) {
return Err(RuntimeSelection::DoesNotServe {
runtime: name.to_string(),
format: format.to_string(),
serves,
});
}
Ok((runtime, device))
}
pub fn names(&self) -> Vec<&'static str> {
self.by_name.keys().copied().collect()
}
pub fn is_empty(&self) -> bool {
self.by_name.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::fixture;
use serde_json::json;
#[test]
fn the_binding_covers_what_a_load_reads_and_nothing_more() {
let base = fixture::manifest();
let fingerprint = |m: &Manifest| LoadBinding::of(m).fingerprint().to_string();
let changed = |f: fn(&mut Manifest)| {
let mut m = base.clone();
f(&mut m);
fingerprint(&m)
};
let untouched = fingerprint(&base);
assert!(
crate::crypto::is_sha256_digest(&untouched),
"{untouched} is the one digest spelling"
);
assert_eq!(untouched, fingerprint(&base.clone()), "and it is stable");
for (what, moved) in [
(
"an input name",
changed(|m| m.inputs[0].name = "boards".into()),
),
(
"an input dtype",
changed(|m| m.inputs[0].dtype = "f64".into()),
),
(
"an input shape",
changed(|m| m.inputs[0].shape = crate::model::manifest::fixed_shape(&[1, 2, 6, 8])),
),
(
"an input axis becoming variable",
changed(|m| {
m.inputs[0].shape[3] = crate::model::manifest::Dim::Named("W".to_string());
}),
),
(
"the name a variable axis is given",
changed(|m| {
m.inputs[0].shape[3] = crate::model::manifest::Dim::Named("H".to_string());
}),
),
(
"an output name",
changed(|m| m.outputs[0].name = "logits".into()),
),
(
"a second input",
changed(|m| {
let mut extra = m.inputs[0].clone();
extra.name = "other".into();
m.inputs.push(extra);
}),
),
(
"a second output",
changed(|m| {
let mut extra = m.outputs[0].clone();
extra.name = "other".into();
m.outputs.push(extra);
}),
),
] {
assert_ne!(untouched, moved, "{what} is part of the binding");
}
let mut two_inputs = base.clone();
let mut extra = two_inputs.inputs[0].clone();
extra.name = "other".to_string();
two_inputs.inputs.push(extra);
let mut swapped = two_inputs.clone();
swapped.inputs.swap(0, 1);
assert_ne!(
fingerprint(&two_inputs),
fingerprint(&swapped),
"input order is part of the binding"
);
let (order_a, order_b) = fixture::two_out();
assert_ne!(
fingerprint(&order_a),
fingerprint(&order_b),
"output order is part of the binding"
);
for (what, moved) in [
("the model name", changed(|m| m.name = "ada.other".into())),
("the version", changed(|m| m.version = "9.9.9".into())),
(
"the description",
changed(|m| m.description = "another graph entirely".into()),
),
(
"the artifact path",
changed(|m| m.artifact = Some("other.onnx".into())),
),
(
"an input adapter",
changed(|m| m.inputs[0].adapter = Some(json!({"var": "data.other"}))),
),
(
"the result expression",
changed(|m| m.result = Some(json!({"other": {"var": "policy"}}))),
),
(
"an output dtype",
changed(|m| m.outputs[0].dtype = "f64".into()),
),
(
"an output shape",
changed(|m| m.outputs[0].shape = crate::model::manifest::fixed_shape(&[1, 9])),
),
] {
assert_eq!(untouched, moved, "{what} is not part of the binding");
}
}
struct Stub(&'static str, &'static [&'static str]);
impl ModelRuntime for Stub {
fn name(&self) -> &'static str {
self.0
}
fn devices(&self) -> &'static [&'static str] {
TRACT_DEVICES
}
fn formats(&self) -> &'static [&'static str] {
self.1
}
fn load(
&self,
_bytes: &[u8],
_binding: &LoadBinding,
_device: &str,
) -> Result<Arc<dyn LoadedModel>, LoadError> {
Err(LoadError::new("parse", "stub"))
}
}
#[test]
fn every_known_runtime_lists_cpu_and_a_format_and_nothing_else_registers() {
let mut union: Vec<&str> = Vec::new();
for name in NAMES {
let devices = devices_of(name).expect("every known runtime lists its devices");
assert!(devices.contains(&"cpu"), "{name} must run on cpu");
let formats = formats_of(name).expect("every known runtime lists its formats");
assert!(!formats.is_empty(), "{name} must serve a format");
for format in formats {
if !union.contains(format) {
union.push(format);
}
}
assert_eq!(intern(name), Some(*name));
}
union.sort_unstable();
let mut known = KNOWN_FORMATS.to_vec();
known.sort_unstable();
assert_eq!(
known, union,
"KNOWN_FORMATS is the union of formats_of over NAMES"
);
assert!(devices_of("onnxruntime").is_none());
assert!(formats_of("onnxruntime").is_none());
assert!(intern("onnxruntime").is_none());
let mut registry = ModelRuntimes::empty();
assert!(registry.is_empty());
registry
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect("a listed name registers");
let err = registry
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect_err("twice");
assert!(err.contains("registered twice"), "{err}");
let err = registry
.register(Arc::new(Stub("onnxruntime", TRACT_FORMATS)))
.expect_err("unlisted");
assert!(err.contains("model::runtimes::NAMES"), "{err}");
let err = ModelRuntimes::empty()
.register(Arc::new(Stub("tract", &[])))
.expect_err("no format");
assert!(err.contains("serves no format"), "{err}");
assert_eq!(registry.names(), ["tract"]);
assert!(registry.get("tract").is_some());
assert!(registry.get("onnxruntime").is_none());
}
fn refused(
result: Result<(Arc<dyn ModelRuntime>, &str), RuntimeSelection>,
) -> RuntimeSelection {
match result {
Err(selection) => selection,
Ok((runtime, device)) => {
unreachable!("expected a refusal, got {} on {device}", runtime.name())
}
}
}
#[test]
fn selection_follows_the_format_table_and_names_every_refusal() {
let config = ModelsConfig::default();
let registry = ModelRuntimes::builtin(&config);
let (runtime, device) = registry
.default_for(&config, "onnx")
.expect("the default config maps onnx to tract");
assert_eq!(runtime.name(), "tract");
assert_eq!(device, "cpu");
let (runtime, device) = registry
.for_format(&config, "tract", "onnx")
.expect("explicit tract for onnx");
assert_eq!(runtime.name(), "tract");
assert_eq!(device, "cpu");
let mut metal = ModelsConfig::default();
metal.runtimes.get_mut("tract").expect("entry").device = "metal".to_string();
let (_, device) = registry
.default_for(&metal, "onnx")
.expect("device comes from the config");
assert_eq!(device, "metal");
let err = refused(registry.default_for(&config, "nnef"));
assert_eq!(
err,
RuntimeSelection::NoDefault {
format: "nnef".to_string()
}
);
assert_eq!(
err.to_string(),
"no default runtime for format 'nnef' (models.default_runtime.nnef is not set)"
);
let mut ort = ModelsConfig::default();
ort.default_runtime
.insert("onnx".to_string(), "ort".to_string());
let err = refused(registry.default_for(&ort, "onnx"));
assert_eq!(
err,
RuntimeSelection::UnknownRuntime {
runtime: "ort".to_string(),
format: "onnx".to_string()
}
);
assert_eq!(
err.to_string(),
"runtime 'ort' for format 'onnx' is unknown: this build lists tract"
);
let mut disabled = ModelsConfig::default();
disabled.runtimes.get_mut("tract").expect("entry").enabled = false;
let err = refused(ModelRuntimes::builtin(&disabled).default_for(&disabled, "onnx"));
assert_eq!(
err,
RuntimeSelection::NotEnabled {
runtime: "tract".to_string(),
format: "onnx".to_string()
}
);
assert_eq!(
err.to_string(),
"runtime 'tract' for format 'onnx' is not enabled on this node \
(models.runtimes.tract.enabled)"
);
let err = refused(ModelRuntimes::empty().default_for(&config, "onnx"));
assert!(matches!(err, RuntimeSelection::NotEnabled { .. }), "{err}");
let mut by_hand = ModelRuntimes::empty();
by_hand
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect("registers");
let mut absent = ModelsConfig::default();
absent.runtimes.clear();
let err = refused(by_hand.for_format(&absent, "tract", "onnx"));
assert!(matches!(err, RuntimeSelection::NotEnabled { .. }), "{err}");
let err = refused(registry.for_format(&config, "tract", "nnef"));
assert_eq!(
err,
RuntimeSelection::DoesNotServe {
runtime: "tract".to_string(),
format: "nnef".to_string(),
serves: TRACT_FORMATS
}
);
assert_eq!(
err.to_string(),
"runtime 'tract' does not serve format 'nnef'; it serves 'onnx'"
);
}
#[test]
fn builtin_follows_the_config() {
let config = ModelsConfig::default();
let registry = ModelRuntimes::builtin(&config);
assert_eq!(registry.names(), ["tract"]);
let tract = registry.get("tract").expect("registered");
assert_eq!(tract.name(), "tract");
assert!(tract.devices().contains(&"cpu"));
assert_eq!(tract.formats(), formats_of("tract").expect("listed"));
let mut disabled = ModelsConfig::default();
disabled
.runtimes
.get_mut("tract")
.expect("default entry")
.enabled = false;
assert!(ModelRuntimes::builtin(&disabled).is_empty());
let mut absent = ModelsConfig::default();
absent.runtimes.clear();
assert!(ModelRuntimes::builtin(&absent).is_empty());
}
#[test]
fn errors_name_their_stage() {
assert_eq!(
LoadError::new("device", "no metal here").to_string(),
"device: no metal here"
);
assert_eq!(
RunError::new("input", "rank 3 given, rank 4 expected").to_string(),
"input: rank 3 given, rank 4 expected"
);
}
}