use crate::{DataType, FeatureInfo, Features, Model, MultiArray, PredictionError, ShapeConstraint};
use super::{AxisRange, ModelDescription};
#[cfg(test)]
mod tests;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum Dim {
Exactly(usize),
#[cfg_attr(
not(any(
feature = "align",
feature = "face",
feature = "siglip",
feature = "speaker",
feature = "whisper"
)),
allow(dead_code, reason = "no door in this feature set reads an axis back")
)]
AnyFixed,
#[cfg_attr(
not(any(feature = "whisper", feature = "align")),
allow(
dead_code,
reason = "the whisper decoder and the aligner's window are this variant's only producers"
)
)]
AtLeast(usize),
#[cfg_attr(
not(feature = "lid"),
allow(dead_code, reason = "the lid door is this variant's only producer")
)]
Range(AxisRange),
}
impl core::fmt::Display for Dim {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Exactly(size) => write!(f, "{size}"),
Self::AnyFixed => f.write_str("any one non-zero fixed size"),
Self::AtLeast(floor) => write!(f, "any one fixed size, at least {floor}"),
Self::Range(range) => write!(f, "{range}"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct FeatureContract {
name: &'static str,
dtype: DataType,
dims: Vec<Dim>,
}
impl FeatureContract {
pub(crate) const fn new(name: &'static str, dtype: DataType, dims: Vec<Dim>) -> Self {
Self { name, dtype, dims }
}
fn required_verdict(&self) -> ShapeConstraint {
if self.dims.iter().any(|d| matches!(d, Dim::Range(_))) {
ShapeConstraint::Range
} else {
ShapeConstraint::Fixed
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum StateContract {
None,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct LoadContract {
inputs: Vec<FeatureContract>,
outputs: Vec<FeatureContract>,
state: StateContract,
}
impl LoadContract {
pub(crate) const fn new(
inputs: Vec<FeatureContract>,
outputs: Vec<FeatureContract>,
state: StateContract,
) -> Self {
Self {
inputs,
outputs,
state,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct MissingFeature {
feature: &'static str,
}
impl MissingFeature {
const fn new(feature: &'static str) -> Self {
Self { feature }
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct DataTypeMismatch {
feature: &'static str,
expected: DataType,
observed: Option<DataType>,
}
impl DataTypeMismatch {
const fn new(feature: &'static str, expected: DataType, observed: Option<DataType>) -> Self {
Self {
feature,
expected,
observed,
}
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(&self) -> String {
self.expected.as_str().to_string()
}
pub(crate) fn observed(&self) -> String {
self
.observed
.map_or_else(|| "none".to_string(), |d| d.as_str().to_string())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct RankMismatch {
feature: &'static str,
expected: usize,
observed: usize,
}
impl RankMismatch {
const fn new(feature: &'static str, expected: usize, observed: usize) -> Self {
Self {
feature,
expected,
observed,
}
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(&self) -> String {
format!("rank {}", self.expected)
}
pub(crate) fn observed(&self) -> String {
format!("rank {}", self.observed)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct FlexibilityMismatch {
feature: &'static str,
expected: ShapeConstraint,
observed: Option<ShapeConstraint>,
}
impl FlexibilityMismatch {
const fn new(
feature: &'static str,
expected: ShapeConstraint,
observed: Option<ShapeConstraint>,
) -> Self {
Self {
feature,
expected,
observed,
}
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(&self) -> String {
self.expected.to_string()
}
pub(crate) fn observed(&self) -> String {
self
.observed
.map_or_else(|| "none".to_string(), |c| c.to_string())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct AxisMismatch {
feature: &'static str,
axis: usize,
expected: Dim,
observed: Option<AxisRange>,
}
impl AxisMismatch {
const fn new(
feature: &'static str,
axis: usize,
expected: Dim,
observed: Option<AxisRange>,
) -> Self {
Self {
feature,
axis,
expected,
observed,
}
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(&self) -> String {
format!("axis {} {}", self.axis, self.expected)
}
pub(crate) fn observed(&self) -> String {
match self.observed {
Some(range) => format!("axis {} {range}", self.axis),
None => format!("axis {} none", self.axis),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ZeroSizedAxis {
feature: &'static str,
axis: usize,
}
impl ZeroSizedAxis {
const fn new(feature: &'static str, axis: usize) -> Self {
Self { feature, axis }
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(&self) -> String {
format!("axis {} {}", self.axis, Dim::AnyFixed)
}
pub(crate) fn observed(&self) -> String {
format!("axis {} 0", self.axis)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct OptionalOutput {
feature: &'static str,
}
impl OptionalOutput {
const fn new(feature: &'static str) -> Self {
Self { feature }
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct UnsatisfiableInput {
name: String,
}
impl UnsatisfiableInput {
const fn new(name: String) -> Self {
Self { name }
}
pub(crate) fn name(&self) -> &str {
&self.name
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct UnsatisfiableState {
name: String,
}
impl UnsatisfiableState {
const fn new(name: String) -> Self {
Self { name }
}
pub(crate) fn name(&self) -> &str {
&self.name
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub(crate) enum ContractViolation {
#[error("model declares no feature `{}`", .0.feature())]
Missing(MissingFeature),
#[error(
"feature `{}` is {}, and the contract states {}",
.0.feature(), .0.observed(), .0.expected()
)]
DataType(DataTypeMismatch),
#[error(
"feature `{}` has {}, and the contract states {}",
.0.feature(), .0.observed(), .0.expected()
)]
Rank(RankMismatch),
#[error(
"feature `{}` is {}, and the contract's axes require {}",
.0.feature(), .0.observed(), .0.expected()
)]
Flexibility(FlexibilityMismatch),
#[error(
"feature `{}` declares {}, and the contract states {}",
.0.feature(), .0.observed(), .0.expected()
)]
Axis(AxisMismatch),
#[error(
"feature `{}` declares {}, and the contract states {}; a door that READS an \
axis back allocates from it, so an empty one loads clean and computes nothing",
.0.feature(), .0.observed(), .0.expected()
)]
ZeroSizedAxis(ZeroSizedAxis),
#[error(
"model declares the output `{}` OPTIONAL, and the contract names it as one the door reads; \
a prediction that omits it satisfies the model and fails the door",
.0.feature()
)]
OptionalOutput(OptionalOutput),
#[error(
"model declares a required input `{}` the contract does not name, so every \
prediction would fail",
.0.name()
)]
UnsatisfiableInput(UnsatisfiableInput),
#[error(
"model declares the state buffer `{}`, and the contract states none",
.0.name()
)]
UnsatisfiableState(UnsatisfiableState),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct FeatureRendering {
feature: &'static str,
expected: String,
actual: String,
}
impl FeatureRendering {
const fn new(feature: &'static str, expected: String, actual: String) -> Self {
Self {
feature,
expected,
actual,
}
}
pub(crate) const fn feature(&self) -> &'static str {
self.feature
}
pub(crate) fn expected(self) -> String {
self.expected
}
pub(crate) fn actual(self) -> String {
self.actual
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum Rendered {
Feature(FeatureRendering),
UnsatisfiableInput(String),
UnsatisfiableState(String),
}
impl ContractViolation {
pub(crate) fn rendered(self) -> Rendered {
let (feature, expected, actual) = match self {
Self::UnsatisfiableInput(input) => {
return Rendered::UnsatisfiableInput(input.name);
}
Self::UnsatisfiableState(state) => {
return Rendered::UnsatisfiableState(state.name);
}
Self::Missing(missing) => (
missing.feature(),
"a declared feature".to_string(),
"missing".to_string(),
),
Self::DataType(mismatch) => (mismatch.feature(), mismatch.expected(), mismatch.observed()),
Self::Rank(mismatch) => (mismatch.feature(), mismatch.expected(), mismatch.observed()),
Self::Flexibility(mismatch) => (mismatch.feature(), mismatch.expected(), mismatch.observed()),
Self::Axis(mismatch) => (mismatch.feature(), mismatch.expected(), mismatch.observed()),
Self::ZeroSizedAxis(zero) => (zero.feature(), zero.expected(), zero.observed()),
Self::OptionalOutput(output) => (
output.feature(),
"a required output".to_string(),
"optional".to_string(),
),
};
Rendered::Feature(FeatureRendering::new(feature, expected, actual))
}
}
pub(crate) fn check_load_contract(
description: &ModelDescription,
contract: &LoadContract,
) -> Result<(), ContractViolation> {
for feature in &contract.inputs {
check_feature_contract(feature, description.input(feature.name))?;
}
for feature in &contract.outputs {
let declared = description.output(feature.name);
check_feature_contract(feature, declared)?;
if declared.is_some_and(FeatureInfo::is_optional) {
return Err(ContractViolation::OptionalOutput(OptionalOutput::new(
feature.name,
)));
}
}
for declared in description.inputs() {
if !declared.is_optional()
&& !contract
.inputs
.iter()
.any(|feature| feature.name == declared.name())
{
return Err(ContractViolation::UnsatisfiableInput(
UnsatisfiableInput::new(declared.name().to_string()),
));
}
}
match contract.state {
StateContract::None => {
if let Some(state) = description.states().first() {
return Err(ContractViolation::UnsatisfiableState(
UnsatisfiableState::new(state.name().to_string()),
));
}
}
}
Ok(())
}
fn check_feature_contract(
contract: &FeatureContract,
declared: Option<&FeatureInfo>,
) -> Result<(), ContractViolation> {
let name = contract.name;
let Some(declared) = declared else {
return Err(ContractViolation::Missing(MissingFeature::new(name)));
};
if declared.data_type() != Some(contract.dtype) {
return Err(ContractViolation::DataType(DataTypeMismatch::new(
name,
contract.dtype,
declared.data_type(),
)));
}
if declared.shape().len() != contract.dims.len() {
return Err(ContractViolation::Rank(RankMismatch::new(
name,
contract.dims.len(),
declared.shape().len(),
)));
}
let required = contract.required_verdict();
if declared.shape_constraint() != Some(required) {
return Err(ContractViolation::Flexibility(FlexibilityMismatch::new(
name,
required,
declared.shape_constraint(),
)));
}
for (axis, dim) in contract.dims.iter().enumerate() {
let observed = declared.axis_ranges().get(axis).copied();
if *dim == Dim::AnyFixed && observed == Some(AxisRange::new(0, 1)) {
return Err(ContractViolation::ZeroSizedAxis(ZeroSizedAxis::new(
name, axis,
)));
}
let satisfied = match *dim {
Dim::Exactly(size) => observed == Some(AxisRange::new(size, 1)),
Dim::AnyFixed => observed.is_some_and(|range| range.count() == 1),
Dim::AtLeast(floor) => {
observed.is_some_and(|range| range.count() == 1 && range.min() >= floor)
}
Dim::Range(range) => observed == Some(range),
};
if !satisfied {
return Err(ContractViolation::Axis(AxisMismatch::new(
name, axis, *dim, observed,
)));
}
}
Ok(())
}
#[derive(Debug)]
pub(crate) struct Checked {
model: Model,
outputs: Vec<&'static str>,
}
impl Checked {
pub(crate) fn new(model: Model, contract: &LoadContract) -> Result<Self, ContractViolation> {
check_load_contract(model.description(), contract)?;
Ok(Self {
model,
outputs: contract.outputs.iter().map(|output| output.name).collect(),
})
}
#[cfg(any(
feature = "align",
feature = "face",
feature = "siglip",
feature = "speaker",
feature = "whisper"
))]
pub(crate) const fn description(&self) -> &ModelDescription {
self.model.description()
}
pub(crate) fn predict_with(
&self,
inputs: &[(&str, &MultiArray)],
) -> Result<Features, PredictionError> {
self.model.predict_with_outputs(inputs, &self.outputs)
}
}