use onnx_runtime_ir::DataType;
use onnx_runtime_ir::DeviceId;
use std::marker::PhantomData;
use crate::error::{EpError, Result};
#[derive(Clone, Copy, Debug)]
pub struct DevicePtr(pub *const std::ffi::c_void);
#[derive(Clone, Copy, Debug)]
pub struct DevicePtrMut(pub *mut std::ffi::c_void);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ExternalMmapRegion {
pub mapping_id: usize,
pub offset: usize,
pub len: usize,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum TensorBacking {
#[default]
Opaque,
ExternalMmap(ExternalMmapRegion),
}
impl DevicePtr {
pub fn as_ptr<T>(self) -> *const T {
self.0 as *const T
}
pub fn is_null(self) -> bool {
self.0.is_null()
}
}
impl DevicePtrMut {
pub fn as_ptr<T>(self) -> *mut T {
self.0 as *mut T
}
pub fn is_null(self) -> bool {
self.0.is_null()
}
}
fn validate_view(
data_is_null: bool,
dtype: DataType,
shape: &[usize],
strides: &[i64],
byte_offset: usize,
) -> Result<()> {
if data_is_null {
return Err(EpError::InvalidTensorView {
reason: "data pointer is null".into(),
});
}
if shape.len() != strides.len() {
return Err(EpError::InvalidTensorView {
reason: format!(
"rank mismatch: shape has {} dims but strides has {}",
shape.len(),
strides.len()
),
});
}
if dtype == DataType::String {
return Err(EpError::InvalidTensorView {
reason: "String dtype has no fixed-width raw layout".into(),
});
}
let esize = dtype.byte_size();
if esize > 1 && !byte_offset.is_multiple_of(esize) {
return Err(EpError::InvalidTensorView {
reason: format!("byte_offset {byte_offset} is not a multiple of element size {esize}"),
});
}
Ok(())
}
#[derive(Clone, Copy)]
pub struct TensorView<'a> {
pub data: DevicePtr,
pub dtype: DataType,
pub shape: &'a [usize],
pub strides: &'a [i64],
pub byte_offset: usize,
pub device: DeviceId,
pub backing: TensorBacking,
_marker: PhantomData<&'a ()>,
}
impl<'a> TensorView<'a> {
pub fn new(
data: DevicePtr,
dtype: DataType,
shape: &'a [usize],
strides: &'a [i64],
device: DeviceId,
) -> Self {
Self {
data,
dtype,
shape,
strides,
byte_offset: 0,
device,
backing: TensorBacking::Opaque,
_marker: PhantomData,
}
}
pub fn absent(dtype: DataType) -> Self {
Self {
data: DevicePtr(std::ptr::null()),
dtype,
shape: &[],
strides: &[],
byte_offset: 0,
device: DeviceId::cpu(),
backing: TensorBacking::Opaque,
_marker: PhantomData,
}
}
pub fn is_absent(&self) -> bool {
self.data.is_null()
}
pub fn with_byte_offset(mut self, byte_offset: usize) -> Self {
self.byte_offset = byte_offset;
self
}
pub fn with_backing(mut self, backing: TensorBacking) -> Self {
self.backing = backing;
self
}
pub fn validate(&self) -> Result<()> {
validate_view(
self.data.is_null(),
self.dtype,
self.shape,
self.strides,
self.byte_offset,
)
}
pub fn is_contiguous(&self) -> bool {
onnx_runtime_ir::is_contiguous(self.shape, self.strides)
}
pub fn numel(&self) -> usize {
self.shape.iter().product()
}
pub fn byte_size(&self) -> usize {
self.dtype.storage_bytes(self.numel())
}
pub fn data_ptr<T>(&self) -> *const T {
(self.data.0 as *const u8).wrapping_add(self.byte_offset) as *const T
}
}
pub struct TensorMut<'a> {
pub data: DevicePtrMut,
pub dtype: DataType,
pub shape: &'a [usize],
pub strides: &'a [i64],
pub byte_offset: usize,
pub device: DeviceId,
_marker: PhantomData<&'a mut ()>,
}
impl<'a> TensorMut<'a> {
pub fn new(
data: DevicePtrMut,
dtype: DataType,
shape: &'a [usize],
strides: &'a [i64],
device: DeviceId,
) -> Self {
Self {
data,
dtype,
shape,
strides,
byte_offset: 0,
device,
_marker: PhantomData,
}
}
pub fn with_byte_offset(mut self, byte_offset: usize) -> Self {
self.byte_offset = byte_offset;
self
}
pub fn validate(&self) -> Result<()> {
validate_view(
self.data.is_null(),
self.dtype,
self.shape,
self.strides,
self.byte_offset,
)
}
pub fn is_contiguous(&self) -> bool {
onnx_runtime_ir::is_contiguous(self.shape, self.strides)
}
pub fn numel(&self) -> usize {
self.shape.iter().product()
}
pub fn byte_size(&self) -> usize {
self.dtype.storage_bytes(self.numel())
}
pub fn data_ptr_mut<T>(&mut self) -> *mut T {
(self.data.0 as *mut u8).wrapping_add(self.byte_offset) as *mut T
}
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_runtime_ir::compute_contiguous_strides;
fn ptr(buf: &[u8]) -> DevicePtr {
DevicePtr(buf.as_ptr() as *const std::ffi::c_void)
}
#[test]
fn contiguous_view_roundtrips_invariants() {
let buf = vec![0u8; 6 * 4];
let shape = [2usize, 3];
let strides = compute_contiguous_strides(&shape);
let v = TensorView::new(
ptr(&buf),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
v.validate().unwrap();
assert!(v.is_contiguous());
assert_eq!(v.numel(), 6);
assert_eq!(v.byte_size(), 24);
assert_eq!(v.byte_offset, 0);
}
#[test]
fn strided_noncontiguous_view_is_representable() {
let buf = vec![0u8; 6 * 4];
let shape = [3usize, 2];
let strides = [1i64, 3];
let v = TensorView::new(
ptr(&buf),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
)
.with_byte_offset(4);
v.validate().unwrap();
assert!(!v.is_contiguous());
assert_eq!(v.shape, &[3, 2]);
assert_eq!(v.strides, &[1, 3]);
assert_eq!(v.byte_offset, 4);
let base = buf.as_ptr() as usize;
assert_eq!(v.data_ptr::<f32>() as usize, base + 4);
}
#[test]
fn negative_strides_are_representable() {
let buf = vec![0u8; 4 * 4];
let shape = [4usize];
let strides = [-1i64];
let v = TensorView::new(
ptr(&buf),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
v.validate().unwrap();
assert!(!v.is_contiguous());
}
#[test]
fn validate_rejects_rank_mismatch() {
let buf = vec![0u8; 8];
let shape = [2usize, 2];
let strides = [1i64]; let v = TensorView::new(
ptr(&buf),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
assert!(v.validate().is_err());
}
#[test]
fn validate_rejects_misaligned_offset_and_string_and_null() {
let buf = vec![0u8; 16];
let shape = [2usize];
let strides = [1i64];
let bad_off = TensorView::new(
ptr(&buf),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
)
.with_byte_offset(1);
assert!(bad_off.validate().is_err());
let bad_dt = TensorView::new(
ptr(&buf),
DataType::String,
&shape,
&strides,
DeviceId::cpu(),
);
assert!(bad_dt.validate().is_err());
let bad_null = TensorView::new(
DevicePtr(std::ptr::null()),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
assert!(bad_null.validate().is_err());
}
#[test]
fn mut_view_offset_pointer() {
let mut buf = vec![0u8; 4 * 4];
let base = buf.as_ptr() as usize;
let shape = [2usize];
let strides = [1i64];
let mut v = TensorMut::new(
DevicePtrMut(buf.as_mut_ptr() as *mut std::ffi::c_void),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
)
.with_byte_offset(8);
v.validate().unwrap();
assert_eq!(v.data_ptr_mut::<f32>() as usize, base + 8);
}
#[test]
fn sub_byte_byte_size_uses_packing() {
let buf = vec![0u8; 4];
let shape = [5usize];
let strides = [1i64];
let v = TensorView::new(ptr(&buf), DataType::Int4, &shape, &strides, DeviceId::cpu());
assert_eq!(v.byte_size(), 3);
v.validate().unwrap();
}
}