use crate::model::contract::{ContractViolation, Rendered};
pub type Result<T> = core::result::Result<T, Error>;
pub use windit::WinditError;
pub use crate::embeddings::tokenizer_guard::PostProcessorTemplate;
#[derive(Debug)]
pub struct ContractMismatch {
feature: &'static str,
expected: String,
actual: String,
}
impl ContractMismatch {
#[inline(always)]
pub const fn new(feature: &'static str, expected: String, actual: String) -> Self {
Self {
feature,
expected,
actual,
}
}
#[inline(always)]
pub const fn feature(&self) -> &'static str {
self.feature
}
#[inline(always)]
pub fn expected(&self) -> &str {
&self.expected
}
#[inline(always)]
pub fn actual(&self) -> &str {
&self.actual
}
}
#[derive(Debug)]
pub struct OutputShape {
got: Vec<usize>,
expected: Vec<usize>,
}
impl OutputShape {
#[inline(always)]
pub const fn new(got: Vec<usize>, expected: Vec<usize>) -> Self {
Self { got, expected }
}
#[inline(always)]
pub fn got(&self) -> &[usize] {
&self.got
}
#[inline(always)]
pub fn expected(&self) -> &[usize] {
&self.expected
}
}
#[derive(Debug)]
pub struct AudioTooLong {
len: usize,
max: usize,
}
#[allow(clippy::len_without_is_empty)]
impl AudioTooLong {
#[inline(always)]
pub const fn new(len: usize, max: usize) -> Self {
Self { len, max }
}
#[inline(always)]
pub const fn len(&self) -> usize {
self.len
}
#[inline(always)]
pub const fn max(&self) -> usize {
self.max
}
}
#[derive(Debug)]
pub struct EmbeddingDimMismatch {
expected: usize,
got: usize,
}
impl EmbeddingDimMismatch {
#[inline(always)]
pub const fn new(expected: usize, got: usize) -> Self {
Self { expected, got }
}
#[inline(always)]
pub const fn expected(&self) -> usize {
self.expected
}
#[inline(always)]
pub const fn got(&self) -> usize {
self.got
}
}
#[derive(Debug)]
pub struct SpecialTokenOverhead {
added: usize,
window: usize,
}
impl SpecialTokenOverhead {
#[inline(always)]
pub const fn new(added: usize, window: usize) -> Self {
Self { added, window }
}
#[inline(always)]
pub const fn added(&self) -> usize {
self.added
}
#[inline(always)]
pub const fn window(&self) -> usize {
self.window
}
}
#[derive(Debug)]
pub struct TokenCount {
got: usize,
max: usize,
}
impl TokenCount {
#[inline(always)]
pub const fn new(got: usize, max: usize) -> Self {
Self { got, max }
}
#[inline(always)]
pub const fn got(&self) -> usize {
self.got
}
#[inline(always)]
pub const fn max(&self) -> usize {
self.max
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum Error {
#[error("failed to load model: {0}")]
Load(#[from] crate::LoadError),
#[error("prediction failed: {0}")]
Prediction(#[from] crate::PredictionError),
#[error("tensor failed: {0}")]
Tensor(#[from] crate::TensorError),
#[error(
"model contract mismatch on `{}`: expected {}, got {}",
.0.feature(),
.0.expected(),
.0.actual()
)]
ContractMismatch(ContractMismatch),
#[error(
"model declares a required input `{0}` that this door never supplies, so \
every prediction would fail"
)]
UnsatisfiableInput(String),
#[error(
"model declares the state buffer `{0}`, and this door predicts through the \
stateless API; a stateful graph needs an `MLState` on every prediction"
)]
UnsatisfiableState(String),
#[error("output shape mismatch: expected {:?}, got {:?}", .0.expected(), .0.got())]
OutputShape(OutputShape),
#[error("audio input contains a non-finite value at index {0}")]
NonFiniteInput(usize),
#[error("model output contains a non-finite value at index {0}")]
NonFiniteOutput(usize),
#[error("audio input is empty")]
EmptyAudio,
#[error(
"audio window has {} samples, over the {}-sample per-window limit; use \
`AudioEncoder::embed_windows` for long audio",
.0.len(),
.0.max()
)]
AudioTooLong(AudioTooLong),
#[error("text input is empty")]
EmptyText,
#[error("embedding dimension mismatch: expected {}, got {}", .0.expected(), .0.got())]
EmbeddingDimMismatch(EmbeddingDimMismatch),
#[error("embedding contains a non-finite value at component {0}")]
NonFiniteEmbedding(usize),
#[error("embedding has zero magnitude and cannot be normalized")]
EmbeddingZero,
#[error("embedding is not unit-norm: |norm² − 1| = {0}")]
EmbeddingNotUnitNorm(f32),
#[error("failed to load tokenizer: {0}")]
TokenizerLoad(#[source] tokenizers::Error),
#[error("failed to configure tokenizer: {0}")]
TokenizerConfig(#[source] tokenizers::Error),
#[error(
"tokenizer post-processor adds {} special tokens, leaving no room for text in the {}-token window",
.0.added(),
.0.window()
)]
SpecialTokenOverhead(SpecialTokenOverhead),
#[error("tokenizer post-processor is inconsistent: {0}")]
PostProcessorTemplate(PostProcessorTemplate),
#[error("failed to tokenize text: {0}")]
Tokenize(#[source] tokenizers::Error),
#[error("tokenized input has {} tokens, exceeding the fixed {}-token window", .0.got(), .0.max())]
TokenCount(TokenCount),
#[error("token id {0} exceeds the model's int32 input range")]
TokenIdRange(u32),
#[error("cannot aggregate zero window embeddings")]
EmptyWindows,
#[error("windowed-sequence processing failed: {0}")]
Windowing(#[source] WinditError),
}
impl From<WinditError> for Error {
fn from(e: WinditError) -> Self {
Error::Windowing(e)
}
}
#[cfg(test)]
mod tests;
pub(crate) fn contract_violation(violation: ContractViolation) -> Error {
match violation.rendered() {
Rendered::UnsatisfiableInput(name) => Error::UnsatisfiableInput(name),
Rendered::UnsatisfiableState(name) => Error::UnsatisfiableState(name),
Rendered::Feature(feature) => Error::ContractMismatch(ContractMismatch::new(
feature.feature(),
feature.clone().expected(),
feature.actual(),
)),
}
}