use crate::error::InferenceError;
#[derive(Debug)]
pub(crate) struct IngestedTensor<'a> {
source: &'a str,
tensor_name: &'a str,
shape: &'a [usize],
payload: IngestPayload<'a>,
}
#[derive(Debug)]
enum IngestPayload<'a> {
DecodedF32 {
values: &'a [f32],
dtype_label: &'static str,
},
}
impl<'a> IngestedTensor<'a> {
pub(crate) fn decoded_f32(
source: &'a str,
tensor_name: &'a str,
shape: &'a [usize],
dtype_label: &'static str,
values: &'a [f32],
) -> Self {
Self {
source,
tensor_name,
shape,
payload: IngestPayload::DecodedF32 {
values,
dtype_label,
},
}
}
}
pub(crate) fn validate_ingested_tensor(tensor: IngestedTensor<'_>) -> Result<(), InferenceError> {
let numel = tensor
.shape
.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
.ok_or_else(|| {
InferenceError::InvalidSafetensors(format!(
"{}: tensor {} shape {:?} overflows usize element count",
tensor.source, tensor.tensor_name, tensor.shape,
))
})?;
match tensor.payload {
IngestPayload::DecodedF32 {
values,
dtype_label,
} => {
if values.len() != numel {
return Err(InferenceError::InvalidSafetensors(format!(
"{}: tensor {} ({dtype_label}) decoded element count {} does not match \
shape {:?} (expected {numel})",
tensor.source,
tensor.tensor_name,
values.len(),
tensor.shape,
)));
}
if let Some((idx, bad)) = values.iter().enumerate().find(|(_, v)| !v.is_finite()) {
return Err(InferenceError::InvalidSafetensors(format!(
"{}: tensor {} ({dtype_label}) has non-finite value {bad} at element index \
{idx} of {numel} (shape {:?})",
tensor.source, tensor.tensor_name, tensor.shape,
)));
}
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn accepts_finite_values_including_signed_zero_and_subnormal() {
let values = [0.0f32, -0.0, 1.0, -1.0, f32::MIN_POSITIVE / 2.0];
let tensor = IngestedTensor::decoded_f32("test", "t", &[5], "F32", &values);
assert!(validate_ingested_tensor(tensor).is_ok());
}
#[test]
fn rejects_nan() {
let values = [1.0f32, f32::NAN, 3.0];
let tensor = IngestedTensor::decoded_f32("test", "t", &[3], "F32", &values);
let err = validate_ingested_tensor(tensor).expect_err("NaN must be rejected");
let msg = err.to_string();
assert!(msg.contains('t'), "error should name the tensor: {msg}");
assert!(
msg.contains("element index 1"),
"error should point at the offending index: {msg}"
);
}
#[test]
fn rejects_positive_infinity() {
let values = [1.0f32, f32::INFINITY];
let tensor = IngestedTensor::decoded_f32("test", "t", &[2], "F16", &values);
let err = validate_ingested_tensor(tensor).expect_err("+inf must be rejected");
assert!(err.to_string().contains("element index 1"));
}
#[test]
fn rejects_negative_infinity() {
let values = [f32::NEG_INFINITY, 2.0];
let tensor = IngestedTensor::decoded_f32("test", "t", &[2], "BF16", &values);
let err = validate_ingested_tensor(tensor).expect_err("-inf must be rejected");
assert!(err.to_string().contains("element index 0"));
}
#[test]
fn rejects_shape_product_overflow() {
let values: [f32; 0] = [];
let huge_shape = [usize::MAX, 2];
let tensor = IngestedTensor::decoded_f32("test", "t", &huge_shape, "F32", &values);
let err = validate_ingested_tensor(tensor).expect_err("overflow must be rejected");
assert!(err.to_string().contains("overflows"));
}
#[test]
fn rejects_element_count_mismatch() {
let values = [1.0f32, 2.0, 3.0];
let tensor = IngestedTensor::decoded_f32("test", "t", &[4], "F32", &values);
let err = validate_ingested_tensor(tensor).expect_err("count mismatch must be rejected");
assert!(err.to_string().contains("does not match"));
}
#[test]
fn all_of_a_finite_tensor_is_scanned_not_just_the_first_element() {
let mut values = vec![1.0f32; 64];
values[63] = f32::NAN;
let tensor = IngestedTensor::decoded_f32("test", "t", &[64], "F32", &values);
let err = validate_ingested_tensor(tensor).expect_err("tail NaN must be caught");
assert!(err.to_string().contains("element index 63"));
}
}