Skip to main content

tflite_dyn/sys/
tflite.rs

1use std::{
2    ffi::{c_void, CStr},
3    os::raw::c_char,
4    sync::Arc,
5};
6
7use libloading::Library;
8use semver::{Version, VersionReq};
9
10use crate::Error;
11
12pub struct TfLiteVt {
13    pub library: Arc<Library>,
14    pub version: TfLiteVersionF,
15    pub model_create: TfLiteModelCreateF,
16    pub model_delete: TfLiteModelDeleteF,
17    pub interpreter_options_create: TfLiteInterpreterOptionsCreateF,
18    pub interpreter_options_delete: TfLiteInterpreterOptionsDeleteF,
19    pub interpreter_options_set_num_threads: TfLiteInterpreterOptionsSetNumThreadsF,
20    pub interpreter_options_add_delegate: TfLiteInterpreterOptionsAddDelegateF,
21    pub interpreter_create: TfLiteInterpreterCreateF,
22    pub interpreter_delete: TfLiteInterpreterDeleteF,
23    pub interpreter_get_input_tensor_count: TfLiteInterpreterGetInputTensorCountF,
24    pub interpreter_get_input_tensor: TfLiteInterpreterGetInputTensorF,
25    pub interpreter_allocate_tensors: TfLiteInterpreterAllocateTensorsF,
26    pub interpreter_invoke: TfLiteInterpreterInvokeF,
27    pub interpreter_get_output_tensor_count: TfLiteInterpreterGetOutputTensorCountF,
28    pub interpreter_get_output_tensor: TfLiteInterpreterGetOutputTensorF,
29    pub tensor_type: TfLiteTensorTypeF,
30    pub tensor_num_dims: TfLiteTensorNumDimsF,
31    pub tensor_dim: TfLiteTensorDimF,
32    pub tensor_byte_size: TfLiteTensorByteSizeF,
33    pub tensor_data: TfLiteTensorDataF,
34    pub tensor_name: TfLiteTensorNameF,
35}
36
37impl TfLiteVt {
38    pub fn load(library: Arc<Library>) -> Result<Self, Error> {
39        let version: TfLiteVersionF = unsafe { *library.get(b"TfLiteVersion\0").unwrap() };
40
41        // Validate the DLL version is compatible with the header version we're targeting
42        let version_str = unsafe { CStr::from_ptr((version)()) };
43        let dll_version = Version::parse(version_str.to_str().unwrap()).unwrap();
44        let target_version = VersionReq::parse("2.8.0").unwrap();
45        if !target_version.matches(&dll_version) {
46            return Err(Error::FailedToLoad);
47        }
48
49        let model_create = unsafe { *library.get(b"TfLiteModelCreate\0").unwrap() };
50        let model_delete = unsafe { *library.get(b"TfLiteModelDelete\0").unwrap() };
51        let interpreter_options_create =
52            unsafe { *library.get(b"TfLiteInterpreterOptionsCreate\0").unwrap() };
53        let interpreter_options_delete =
54            unsafe { *library.get(b"TfLiteInterpreterOptionsDelete\0").unwrap() };
55        let interpreter_options_set_num_threads = unsafe {
56            *library
57                .get(b"TfLiteInterpreterOptionsSetNumThreads\0")
58                .unwrap()
59        };
60        let interpreter_options_add_delegate = unsafe {
61            *library
62                .get(b"TfLiteInterpreterOptionsAddDelegate\0")
63                .unwrap()
64        };
65        let interpreter_create = unsafe { *library.get(b"TfLiteInterpreterCreate\0").unwrap() };
66        let interpreter_delete = unsafe { *library.get(b"TfLiteInterpreterDelete\0").unwrap() };
67        let interpreter_get_input_tensor_count = unsafe {
68            *library
69                .get(b"TfLiteInterpreterGetInputTensorCount\0")
70                .unwrap()
71        };
72        let interpreter_get_input_tensor =
73            unsafe { *library.get(b"TfLiteInterpreterGetInputTensor\0").unwrap() };
74        let interpreter_allocate_tensors =
75            unsafe { *library.get(b"TfLiteInterpreterAllocateTensors\0").unwrap() };
76        let interpreter_invoke = unsafe { *library.get(b"TfLiteInterpreterInvoke\0").unwrap() };
77        let interpreter_get_output_tensor_count = unsafe {
78            *library
79                .get(b"TfLiteInterpreterGetOutputTensorCount\0")
80                .unwrap()
81        };
82        let interpreter_get_output_tensor =
83            unsafe { *library.get(b"TfLiteInterpreterGetOutputTensor\0").unwrap() };
84        let tensor_type = unsafe { *library.get(b"TfLiteTensorType\0").unwrap() };
85        let tensor_num_dims = unsafe { *library.get(b"TfLiteTensorNumDims\0").unwrap() };
86        let tensor_dim = unsafe { *library.get(b"TfLiteTensorDim\0").unwrap() };
87        let tensor_byte_size = unsafe { *library.get(b"TfLiteTensorByteSize\0").unwrap() };
88        let tensor_data = unsafe { *library.get(b"TfLiteTensorData\0").unwrap() };
89        let tensor_name = unsafe { *library.get(b"TfLiteTensorName\0").unwrap() };
90
91        Ok(Self {
92            library,
93            version,
94            model_create,
95            model_delete,
96            interpreter_options_create,
97            interpreter_options_delete,
98            interpreter_options_set_num_threads,
99            interpreter_options_add_delegate,
100            interpreter_create,
101            interpreter_delete,
102            interpreter_get_input_tensor_count,
103            interpreter_get_input_tensor,
104            interpreter_allocate_tensors,
105            interpreter_invoke,
106            interpreter_get_output_tensor_count,
107            interpreter_get_output_tensor,
108            tensor_type,
109            tensor_num_dims,
110            tensor_dim,
111            tensor_byte_size,
112            tensor_data,
113            tensor_name,
114        })
115    }
116}
117
118pub type TfLiteVersionF = unsafe extern "C" fn() -> *const c_char;
119
120pub type TfLiteModelCreateF =
121    unsafe extern "C" fn(model_data: *const c_void, size: usize) -> *mut TfLiteModel;
122
123pub type TfLiteModelDeleteF = unsafe extern "C" fn(model: *mut TfLiteModel);
124
125pub type TfLiteInterpreterOptionsCreateF = unsafe extern "C" fn() -> *mut TfLiteInterpreterOptions;
126
127pub type TfLiteInterpreterOptionsDeleteF =
128    unsafe extern "C" fn(options: *mut TfLiteInterpreterOptions);
129
130pub type TfLiteInterpreterOptionsSetNumThreadsF =
131    unsafe extern "C" fn(options: *mut TfLiteInterpreterOptions, num_threads: i32);
132
133pub type TfLiteInterpreterOptionsAddDelegateF =
134    unsafe extern "C" fn(options: *mut TfLiteInterpreterOptions, delegate: *mut TfLiteDelegate);
135
136pub type TfLiteInterpreterCreateF = unsafe extern "C" fn(
137    model: *const TfLiteModel,
138    options: *const TfLiteInterpreterOptions,
139) -> *mut TfLiteInterpreter;
140
141pub type TfLiteInterpreterDeleteF = unsafe extern "C" fn(interpreter: *mut TfLiteInterpreter);
142
143pub type TfLiteInterpreterGetInputTensorCountF =
144    unsafe extern "C" fn(interpreter: *const TfLiteInterpreter) -> i32;
145
146pub type TfLiteInterpreterGetInputTensorF =
147    unsafe extern "C" fn(interpreter: *const TfLiteInterpreter, index: i32) -> *mut TfLiteTensor;
148
149pub type TfLiteInterpreterAllocateTensorsF =
150    unsafe extern "C" fn(interpreter: *mut TfLiteInterpreter) -> TfLiteStatus;
151
152pub type TfLiteInterpreterInvokeF =
153    unsafe extern "C" fn(interpreter: *mut TfLiteInterpreter) -> TfLiteStatus;
154
155pub type TfLiteInterpreterGetOutputTensorCountF =
156    unsafe extern "C" fn(interpreter: *const TfLiteInterpreter) -> i32;
157
158pub type TfLiteInterpreterGetOutputTensorF =
159    unsafe extern "C" fn(interpreter: *const TfLiteInterpreter, index: i32) -> *mut TfLiteTensor;
160
161pub type TfLiteTensorTypeF = unsafe extern "C" fn(tensor: *const TfLiteTensor) -> TfLiteType;
162
163pub type TfLiteTensorNumDimsF = unsafe extern "C" fn(tensor: *const TfLiteTensor) -> i32;
164
165pub type TfLiteTensorDimF =
166    unsafe extern "C" fn(tensor: *const TfLiteTensor, dim_index: i32) -> i32;
167
168pub type TfLiteTensorByteSizeF = unsafe extern "C" fn(tensor: *const TfLiteTensor) -> usize;
169
170pub type TfLiteTensorDataF = unsafe extern "C" fn(tensor: *const TfLiteTensor) -> *const c_void;
171
172pub type TfLiteTensorNameF = unsafe extern "C" fn(tensor: *const TfLiteTensor) -> *const c_char;
173
174#[repr(C)]
175pub struct TfLiteModel {
176    private: [u8; 0],
177}
178
179#[repr(C)]
180pub struct TfLiteInterpreterOptions {
181    private: [u8; 0],
182}
183
184#[repr(C)]
185pub struct TfLiteDelegate {
186    private: [u8; 0],
187}
188
189#[repr(C)]
190pub struct TfLiteInterpreter {
191    private: [u8; 0],
192}
193
194#[repr(C)]
195pub struct TfLiteTensor {
196    private: [u8; 0],
197}
198
199#[repr(C)]
200#[derive(PartialEq, Eq, Debug, Copy, Clone)]
201pub enum TfLiteStatus {
202    Ok = 0,
203    Error = 1,
204    DelegateError = 2,
205    ApplicationError = 3,
206    DelegateDataNotFound = 4,
207    DelegateDataWriteError = 5,
208    DelegateDataReadError = 6,
209    UnresolvedOps = 7,
210}
211
212#[repr(C)]
213#[derive(PartialEq, Eq, Debug, Copy, Clone)]
214pub enum TfLiteType {
215    NoType = 0,
216    Float32 = 1,
217    Int32 = 2,
218    UInt8 = 3,
219    Int64 = 4,
220    String = 5,
221    Bool = 6,
222    Int16 = 7,
223    Complex64 = 8,
224    Int8 = 9,
225    Float16 = 10,
226    Float64 = 11,
227    Complex128 = 12,
228    UInt64 = 13,
229    Resource = 14,
230    Variant = 15,
231    UInt32 = 16,
232}