use std::ffi::{c_void, CStr};
use crate::bindings;
use crate::bindings::*;
use crate::{Error, ErrorKind, Result};
use std::fmt::{Debug, Formatter};
use std::marker::PhantomData;
#[derive(Copy, Clone, PartialEq, Debug, PartialOrd)]
pub struct QuantizationParameters {
pub scale: f32,
pub zero_point: i32,
}
#[derive(Copy, Clone, Eq, PartialEq, Debug, Hash)]
pub enum DataType {
Bool,
UInt8,
UInt16,
UInt32,
UInt64,
Int4,
Int8,
Int16,
Int32,
Int64,
BFloat16,
Float16,
Float32,
Float64,
String,
Complex64,
Complex128,
}
impl DataType {
pub(crate) fn new(tflite_type: TfLiteType) -> Option<DataType> {
match tflite_type {
bindings::TfLiteType_kTfLiteBool => Some(DataType::Bool),
bindings::TfLiteType_kTfLiteUInt8 => Some(DataType::UInt8),
bindings::TfLiteType_kTfLiteUInt16 => Some(DataType::UInt16),
bindings::TfLiteType_kTfLiteUInt32 => Some(DataType::UInt32),
bindings::TfLiteType_kTfLiteUInt64 => Some(DataType::UInt64),
bindings::TfLiteType_kTfLiteInt4 => Some(DataType::Int4),
bindings::TfLiteType_kTfLiteInt8 => Some(DataType::Int8),
bindings::TfLiteType_kTfLiteInt16 => Some(DataType::Int16),
bindings::TfLiteType_kTfLiteInt32 => Some(DataType::Int32),
bindings::TfLiteType_kTfLiteInt64 => Some(DataType::Int64),
bindings::TfLiteType_kTfLiteBFloat16 => Some(DataType::BFloat16),
bindings::TfLiteType_kTfLiteFloat16 => Some(DataType::Float16),
bindings::TfLiteType_kTfLiteFloat32 => Some(DataType::Float32),
bindings::TfLiteType_kTfLiteFloat64 => Some(DataType::Float64),
bindings::TfLiteType_kTfLiteString => Some(DataType::String),
bindings::TfLiteType_kTfLiteComplex64 => Some(DataType::Complex64),
bindings::TfLiteType_kTfLiteComplex128 => Some(DataType::Complex128),
_ => None,
}
}
}
#[derive(Clone, Eq, PartialEq, Debug, Hash)]
pub struct Shape {
rank: usize,
dimensions: Vec<usize>,
}
impl Shape {
pub fn new(dimensions: Vec<usize>) -> Shape {
Shape {
rank: dimensions.len(),
dimensions,
}
}
pub fn dimensions(&self) -> &Vec<usize> {
&self.dimensions
}
pub fn rank(&self) -> usize {
self.rank
}
}
pub(crate) struct TensorData {
data_ptr: *mut u8,
data_length: usize,
}
pub struct Tensor<'a> {
name: String,
data_type: DataType,
shape: Shape,
data: TensorData,
quantization_parameters: Option<QuantizationParameters>,
tensor_ptr: *mut TfLiteTensor,
phantom: PhantomData<&'a TfLiteTensor>,
}
impl Debug for Tensor<'_> {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Tensor")
.field("name", &self.name)
.field("shape", &self.shape)
.field("data_type", &self.data_type)
.field("quantization_parameters", &self.quantization_parameters)
.finish()
}
}
impl<'a> Tensor<'a> {
pub(crate) fn from_raw(tensor_ptr: *mut TfLiteTensor) -> Result<Tensor<'a>> {
unsafe {
if tensor_ptr.is_null() {
return Err(Error::new(ErrorKind::ReadTensorError));
}
let name_ptr = TfLiteTensorName(tensor_ptr);
if name_ptr.is_null() {
return Err(Error::new(ErrorKind::ReadTensorError));
}
let data_ptr = TfLiteTensorData(tensor_ptr) as *mut u8;
if data_ptr.is_null() {
return Err(Error::new(ErrorKind::ReadTensorError));
}
let name = CStr::from_ptr(name_ptr).to_str().unwrap().to_owned();
let data_length = TfLiteTensorByteSize(tensor_ptr);
let data_type = DataType::new(TfLiteTensorType(tensor_ptr))
.ok_or_else(|| Error::new(ErrorKind::InvalidTensorDataType))?;
let rank = TfLiteTensorNumDims(tensor_ptr);
let dimensions = (0..rank)
.map(|i| TfLiteTensorDim(tensor_ptr, i) as usize)
.collect();
let shape = Shape::new(dimensions);
let data = TensorData {
data_ptr,
data_length,
};
let quantization_parameters_ptr = TfLiteTensorQuantizationParams(tensor_ptr);
let scale = quantization_parameters_ptr.scale;
let quantization_parameters =
if scale == 0.0 || (data_type != DataType::UInt8 && data_type != DataType::Int8) {
None
} else {
Some(QuantizationParameters {
scale: quantization_parameters_ptr.scale,
zero_point: quantization_parameters_ptr.zero_point,
})
};
Ok(Tensor {
name,
data_type,
shape,
data,
quantization_parameters,
tensor_ptr,
phantom: PhantomData,
})
}
}
pub fn shape(&self) -> &Shape {
&self.shape
}
pub fn data<T>(&self) -> &[T] {
let element_size = std::mem::size_of::<T>();
if self.data.data_length % element_size != 0 {
panic!(
"data length {} should be divisible by size of type {}",
self.data.data_length, element_size
)
}
unsafe {
std::slice::from_raw_parts(
self.data.data_ptr as *const T,
self.data.data_length / element_size,
)
}
}
pub fn set_data<T>(&self, data: &[T]) -> Result<()> {
let input_byte_count = std::mem::size_of_val(data);
if self.data.data_length != input_byte_count {
return Err(Error::new(ErrorKind::InvalidTensorDataCount(
data.len(),
input_byte_count,
)));
}
let status = unsafe {
TfLiteTensorCopyFromBuffer(
self.tensor_ptr,
data.as_ptr() as *const c_void,
input_byte_count,
)
};
if status != TfLiteStatus_kTfLiteOk {
Err(Error::new(ErrorKind::FailedToCopyDataToInputTensor))
} else {
Ok(())
}
}
pub fn data_type(&self) -> DataType {
self.data_type
}
pub fn quantization_parameters(&self) -> Option<QuantizationParameters> {
self.quantization_parameters
}
pub fn name(&self) -> &str {
self.name.as_str()
}
}