Skip to main content

ort_web/
tensor.rs

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	/// Data is stored in WASM linear memory and can be immediately accessed.
21	RustView { ptr: *mut c_void, byte_len: usize },
22	/// Data is stored outside of WASM linear memory (i.e. session output, or a tensor created from anything other than
23	/// a Rust slice) and would need to be retrieved if we try to extract this tensor.
24	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				// cast to some kind of typed array first, then convert to uint8array so we can properly copy
61				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					// we have a download function, but no upload...
86					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	/// Synchronize tensor data from the device/runtime so that it is accessible to Rust code.
225	Rust,
226	/// Synchronize tensor data from Rust code so that it is accessible to the runtime.
227	Runtime
228}
229
230pub trait ValueExt {
231	private_trait!();
232
233	/// Synchronize data between Rust & the runtime.
234	///
235	/// See the [top-level documentation][crate] for more information on synchronization.
236	#[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		// definitely safe regardless of what backend is used since it's highly improbable that a backend's tensor would be
246		// smaller than 4 bytes (which is pointer size on wasm32)
247		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}