1use alloc::{boxed::Box, vec::Vec};
2use core::{ffi::c_void, ptr::NonNull, slice};
3
4use js_sys::Uint8Array;
5use ort::{AsPointer, value::ValueTypeMarker};
6use wasm_bindgen::{JsCast, JsValue};
7use web_sys::{HtmlImageElement, ImageBitmap, ImageData};
8
9pub use crate::binding::{ImageFormat, ImageNorm, ImageTensorLayout};
10use crate::{
11 Error,
12 binding::{self, DataType, ImageDataType},
13 memory::MemoryInfo,
14 util::num_elements
15};
16
17pub const TENSOR_SENTINEL: [u8; 4] = [0xFC, 0x86, 0xA5, 0x39];
18
19pub enum TensorData {
20 RustView { ptr: *mut c_void, byte_len: usize },
22 External { buffer: Option<Box<[u8]>> }
25}
26
27#[repr(C)]
28pub struct Tensor {
29 sentinel: [u8; 4],
30 pub js: binding::Tensor,
31 pub data: TensorData,
32 pub memory_info: MemoryInfo
33}
34
35impl Tensor {
36 pub unsafe fn from_ptr(dtype: binding::DataType, ptr: *mut c_void, byte_len: usize, dims: &[i32]) -> Result<Self, JsValue> {
37 let tensor = binding::Tensor::new_from_buffer(dtype, unsafe { buffer_from_ptr(dtype, ptr, byte_len) }, dims)?;
38 Ok(Self {
39 sentinel: TENSOR_SENTINEL,
40 memory_info: MemoryInfo { location: tensor.location() },
41 js: tensor,
42 data: TensorData::RustView { ptr, byte_len }
43 })
44 }
45
46 pub fn from_tensor(tensor: binding::Tensor) -> Self {
47 Self {
48 sentinel: TENSOR_SENTINEL,
49 memory_info: MemoryInfo { location: tensor.location() },
50 js: tensor,
51 data: TensorData::External { buffer: None }
52 }
53 }
54
55 pub async fn sync(&mut self, direction: SyncDirection) -> crate::Result<()> {
56 match direction {
57 SyncDirection::Rust => {
58 let data = self.js.get_data().await?;
59
60 let generic_typed_array = Uint8Array::unchecked_from_js(data);
62 let bytes = Uint8Array::new_with_byte_offset_and_length(
63 &generic_typed_array.buffer(),
64 generic_typed_array.byte_offset(),
65 generic_typed_array.byte_length()
66 );
67 match &mut self.data {
68 TensorData::RustView { ptr, byte_len } => {
69 bytes.copy_to(unsafe { core::slice::from_raw_parts_mut(ptr.cast(), *byte_len) });
70 }
71 TensorData::External { buffer } => {
72 let buffer = match buffer {
73 Some(buffer) => buffer,
74 None => {
75 *buffer = Some(vec![0; generic_typed_array.byte_length() as usize].into_boxed_slice());
76 unsafe { buffer.as_mut().unwrap_unchecked() }
77 }
78 };
79 bytes.copy_to(buffer);
80 }
81 }
82 }
83 SyncDirection::Runtime => {
84 let Ok(generic_typed_array) = self.js.data().map(Uint8Array::unchecked_from_js) else {
85 return Err(Error::new(
87 "Cannot synchronize Rust data to a runtime tensor that is not on the CPU; modify the WebGPU/WebGL buffer directly."
88 ));
89 };
90 let bytes = Uint8Array::new_with_byte_offset_and_length(
91 &generic_typed_array.buffer(),
92 generic_typed_array.byte_offset(),
93 generic_typed_array.byte_length()
94 );
95 bytes.copy_from(match &self.data {
96 TensorData::RustView { ptr, byte_len } => unsafe { core::slice::from_raw_parts(ptr.cast(), *byte_len) },
97 TensorData::External { buffer } => {
98 let Some(buffer) = buffer else {
99 return Ok(());
100 };
101 &*buffer
102 }
103 });
104 }
105 }
106 Ok(())
107 }
108}
109
110pub fn create_buffer(dtype: binding::DataType, shape: &[i32]) -> JsValue {
111 let numel = num_elements(shape) as u32;
112 match dtype {
113 binding::DataType::Bool | binding::DataType::Uint8 => js_sys::Uint8Array::new_with_length(numel).into(),
114 binding::DataType::Int8 => js_sys::Int8Array::new_with_length(numel).into(),
115 binding::DataType::Uint16 => js_sys::Uint16Array::new_with_length(numel).into(),
116 binding::DataType::Int16 => js_sys::Int16Array::new_with_length(numel).into(),
117 binding::DataType::Uint32 => js_sys::Uint32Array::new_with_length(numel).into(),
118 binding::DataType::Int32 => js_sys::Int32Array::new_with_length(numel).into(),
119 binding::DataType::Uint64 => js_sys::BigUint64Array::new_with_length(numel).into(),
120 binding::DataType::Int64 => js_sys::BigInt64Array::new_with_length(numel).into(),
121 binding::DataType::Float32 => js_sys::Float32Array::new_with_length(numel).into(),
122 binding::DataType::Float64 => js_sys::Float64Array::new_with_length(numel).into(),
123 binding::DataType::Int4 | binding::DataType::Uint4 | binding::DataType::Float16 | binding::DataType::String => unimplemented!(),
124 binding::DataType::__Invalid => unreachable!()
125 }
126}
127
128pub unsafe fn buffer_from_ptr(dtype: binding::DataType, ptr: *mut c_void, byte_len: usize) -> JsValue {
129 match dtype {
130 binding::DataType::Bool | binding::DataType::Uint8 => unsafe { js_sys::Uint8Array::view(slice::from_raw_parts(ptr.cast(), byte_len)) }.into(),
131 binding::DataType::Int8 => unsafe { js_sys::Int8Array::view(slice::from_raw_parts(ptr.cast(), byte_len)) }.into(),
132 binding::DataType::Uint16 => unsafe { js_sys::Uint16Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 2)) }.into(),
133 binding::DataType::Int16 => unsafe { js_sys::Int16Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 2)) }.into(),
134 binding::DataType::Uint32 => unsafe { js_sys::Uint32Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 4)) }.into(),
135 binding::DataType::Int32 => unsafe { js_sys::Int32Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 4)) }.into(),
136 binding::DataType::Uint64 => unsafe { js_sys::BigUint64Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 8)) }.into(),
137 binding::DataType::Int64 => unsafe { js_sys::BigInt64Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 8)) }.into(),
138 binding::DataType::Float32 => unsafe { js_sys::Float32Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 4)) }.into(),
139 binding::DataType::Float64 => unsafe { js_sys::Float64Array::view(slice::from_raw_parts(ptr.cast(), byte_len / 8)) }.into(),
140 binding::DataType::Int4 | binding::DataType::Uint4 | binding::DataType::Float16 | binding::DataType::String => unimplemented!(),
141 binding::DataType::__Invalid => unreachable!()
142 }
143}
144
145pub fn dtype_to_onnx(dtype: binding::DataType) -> ort_sys::ONNXTensorElementDataType {
146 match dtype {
147 binding::DataType::String => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING,
148 binding::DataType::Bool => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL,
149 binding::DataType::Uint8 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8,
150 binding::DataType::Int8 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8,
151 binding::DataType::Uint16 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16,
152 binding::DataType::Int16 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16,
153 binding::DataType::Uint32 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32,
154 binding::DataType::Int32 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32,
155 binding::DataType::Uint64 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64,
156 binding::DataType::Int64 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64,
157 binding::DataType::Float16 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16,
158 binding::DataType::Float32 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
159 binding::DataType::Float64 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE,
160 binding::DataType::Int4 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT4,
161 binding::DataType::Uint4 => ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT4,
162 binding::DataType::__Invalid => unreachable!()
163 }
164}
165
166pub fn onnx_to_dtype(dtype: ort_sys::ONNXTensorElementDataType) -> Option<binding::DataType> {
167 match dtype {
168 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING => Some(binding::DataType::String),
169 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL => Some(binding::DataType::Bool),
170 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8 => Some(binding::DataType::Uint8),
171 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8 => Some(binding::DataType::Int8),
172 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16 => Some(binding::DataType::Uint16),
173 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16 => Some(binding::DataType::Int16),
174 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32 => Some(binding::DataType::Uint32),
175 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32 => Some(binding::DataType::Int32),
176 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64 => Some(binding::DataType::Uint64),
177 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64 => Some(binding::DataType::Int64),
178 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16 => Some(binding::DataType::Float16),
179 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT => Some(binding::DataType::Float32),
180 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE => Some(binding::DataType::Float64),
181 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT4 => Some(binding::DataType::Int4),
182 ort_sys::ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT4 => Some(binding::DataType::Uint4),
183 _ => None
184 }
185}
186
187pub struct TypeInfo {
188 pub dtype: ort_sys::ONNXTensorElementDataType,
189 pub shape: Vec<i32>
190}
191
192impl TypeInfo {
193 pub fn new_sys_from_tensor(tensor: &Tensor) -> *mut ort_sys::OrtTypeInfo {
194 Self::new_sys(tensor.js.dtype(), tensor.js.dims())
195 }
196
197 pub fn new_sys_from_value_metadata(metadata: &binding::ValueMetadata) -> *mut ort_sys::OrtTypeInfo {
198 Self::new_sys(
199 metadata.r#type.unwrap(),
200 metadata
201 .shape
202 .as_ref()
203 .unwrap()
204 .iter()
205 .map(|el| match el {
206 binding::ShapeElement::Value(v) => *v as i32,
207 binding::ShapeElement::Named(_) => -1
208 })
209 .collect()
210 )
211 }
212
213 pub fn new_sys(dtype: DataType, shape: Vec<i32>) -> *mut ort_sys::OrtTypeInfo {
214 (Box::leak(Box::new(Self { dtype: dtype_to_onnx(dtype), shape })) as *mut TypeInfo).cast()
215 }
216
217 pub unsafe fn consume_sys(ptr: *mut ort_sys::OrtTypeInfo) -> Box<TypeInfo> {
218 unsafe { Box::from_raw(ptr.cast::<TypeInfo>()) }
219 }
220}
221
222#[derive(Debug, Clone, Copy, PartialEq, Eq)]
223pub enum SyncDirection {
224 Rust,
226 Runtime
228}
229
230pub trait ValueExt {
231 private_trait!();
232
233 #[allow(async_fn_in_trait)]
237 async fn sync(&mut self, direction: SyncDirection) -> crate::Result<()>;
238}
239
240impl<T: ValueTypeMarker> ValueExt for ort::value::Value<T> {
241 private_impl!();
242
243 async fn sync(&mut self, direction: SyncDirection) -> crate::Result<()> {
244 let ptr = self.ptr_mut();
245 let sentinel: [u8; 4] = unsafe { core::ptr::read(ptr.cast()) };
248 if sentinel != TENSOR_SENTINEL {
249 return Err(Error::new("Cannot synchronize Value that was not created by ort-web"));
250 }
251
252 let tensor: &mut Tensor = unsafe { &mut *ptr.cast() };
253 tensor.sync(direction).await
254 }
255}
256
257#[derive(Default)]
258pub struct TensorFromImageOptions {
259 pub norm: Option<ImageNorm>,
260 pub resized_height: Option<u32>,
261 pub resized_width: Option<u32>,
262 pub tensor_format: Option<ImageFormat>,
263 pub tensor_layout: Option<ImageTensorLayout>
264}
265
266#[derive(Default)]
267pub struct TensorFromUrlOptions {
268 pub norm: Option<ImageNorm>,
269 pub resized_height: Option<u32>,
270 pub resized_width: Option<u32>,
271 pub tensor_format: Option<ImageFormat>,
272 pub tensor_layout: Option<ImageTensorLayout>
273}
274
275#[allow(async_fn_in_trait)]
276pub trait TensorFromImage: Sized {
277 private_trait!();
278
279 async fn from_image_data(image_data: &ImageData, options: TensorFromImageOptions) -> crate::Result<Self>;
280 async fn from_image_element(image_element: &HtmlImageElement, options: TensorFromImageOptions) -> crate::Result<Self>;
281 async fn from_image_bitmap(image_bitmap: &ImageBitmap, options: TensorFromImageOptions) -> crate::Result<Self>;
282 async fn from_image_url(url: &str, options: TensorFromImageOptions, original_dimensions: Option<(u32, u32)>) -> crate::Result<Self>;
283}
284
285impl TensorFromImage for ort::value::Tensor<f32> {
286 private_impl!();
287
288 async fn from_image_data(image: &ImageData, options: TensorFromImageOptions) -> crate::Result<Self> {
289 let tensor = Tensor::from_tensor(
290 binding::Tensor::from_image_data(
291 image,
292 &binding::TensorFromImageOptions {
293 data_type: Some(ImageDataType::Float32),
294 norm: options.norm,
295 resized_height: options.resized_height,
296 resized_width: options.resized_width,
297 tensor_format: options.tensor_format,
298 tensor_layout: options.tensor_layout
299 }
300 )
301 .await?
302 );
303 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
304 }
305
306 async fn from_image_element(image: &HtmlImageElement, options: TensorFromImageOptions) -> crate::Result<Self> {
307 let tensor = Tensor::from_tensor(
308 binding::Tensor::from_image_element(
309 image,
310 &binding::TensorFromImageOptions {
311 data_type: Some(ImageDataType::Float32),
312 norm: options.norm,
313 resized_height: options.resized_height,
314 resized_width: options.resized_width,
315 tensor_format: options.tensor_format,
316 tensor_layout: options.tensor_layout
317 }
318 )
319 .await?
320 );
321 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
322 }
323
324 async fn from_image_bitmap(image: &ImageBitmap, options: TensorFromImageOptions) -> crate::Result<Self> {
325 let tensor = Tensor::from_tensor(
326 binding::Tensor::from_image_bitmap(
327 image,
328 &binding::TensorFromImageOptions {
329 data_type: Some(ImageDataType::Float32),
330 norm: options.norm,
331 resized_height: options.resized_height,
332 resized_width: options.resized_width,
333 tensor_format: options.tensor_format,
334 tensor_layout: options.tensor_layout
335 }
336 )
337 .await?
338 );
339 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
340 }
341
342 async fn from_image_url(url: &str, options: TensorFromImageOptions, original_dimensions: Option<(u32, u32)>) -> crate::Result<Self> {
343 let tensor = Tensor::from_tensor(
344 binding::Tensor::from_image_url(
345 url,
346 &binding::TensorFromUrlOptions {
347 base: binding::TensorFromImageOptions {
348 data_type: Some(ImageDataType::Float32),
349 norm: options.norm,
350 resized_height: options.resized_height,
351 resized_width: options.resized_width,
352 tensor_format: options.tensor_format,
353 tensor_layout: options.tensor_layout
354 },
355 width: original_dimensions.map(|(w, _)| w),
356 height: original_dimensions.map(|(_, h)| h)
357 }
358 )
359 .await?
360 );
361 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
362 }
363}
364
365impl TensorFromImage for ort::value::Tensor<u8> {
366 private_impl!();
367
368 async fn from_image_data(image: &ImageData, options: TensorFromImageOptions) -> crate::Result<Self> {
369 let tensor = Tensor::from_tensor(
370 binding::Tensor::from_image_data(
371 image,
372 &binding::TensorFromImageOptions {
373 data_type: Some(ImageDataType::Uint8),
374 norm: options.norm,
375 resized_height: options.resized_height,
376 resized_width: options.resized_width,
377 tensor_format: options.tensor_format,
378 tensor_layout: options.tensor_layout
379 }
380 )
381 .await?
382 );
383 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
384 }
385
386 async fn from_image_element(image: &HtmlImageElement, options: TensorFromImageOptions) -> crate::Result<Self> {
387 let tensor = Tensor::from_tensor(
388 binding::Tensor::from_image_element(
389 image,
390 &binding::TensorFromImageOptions {
391 data_type: Some(ImageDataType::Uint8),
392 norm: options.norm,
393 resized_height: options.resized_height,
394 resized_width: options.resized_width,
395 tensor_format: options.tensor_format,
396 tensor_layout: options.tensor_layout
397 }
398 )
399 .await?
400 );
401 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
402 }
403
404 async fn from_image_bitmap(image: &ImageBitmap, options: TensorFromImageOptions) -> crate::Result<Self> {
405 let tensor = Tensor::from_tensor(
406 binding::Tensor::from_image_bitmap(
407 image,
408 &binding::TensorFromImageOptions {
409 data_type: Some(ImageDataType::Uint8),
410 norm: options.norm,
411 resized_height: options.resized_height,
412 resized_width: options.resized_width,
413 tensor_format: options.tensor_format,
414 tensor_layout: options.tensor_layout
415 }
416 )
417 .await?
418 );
419 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
420 }
421
422 async fn from_image_url(url: &str, options: TensorFromImageOptions, original_dimensions: Option<(u32, u32)>) -> crate::Result<Self> {
423 let tensor = Tensor::from_tensor(
424 binding::Tensor::from_image_url(
425 url,
426 &binding::TensorFromUrlOptions {
427 base: binding::TensorFromImageOptions {
428 data_type: Some(ImageDataType::Uint8),
429 norm: options.norm,
430 resized_height: options.resized_height,
431 resized_width: options.resized_width,
432 tensor_format: options.tensor_format,
433 tensor_layout: options.tensor_layout
434 },
435 width: original_dimensions.map(|(w, _)| w),
436 height: original_dimensions.map(|(_, h)| h)
437 }
438 )
439 .await?
440 );
441 Ok(unsafe { ort::value::Tensor::from_ptr(NonNull::from_mut(Box::leak(Box::new(tensor))).cast(), None) })
442 }
443}