Skip to main content

tenferro_runtime/runtime/
signature.rs

1use num_complex::{Complex32, Complex64};
2use std::mem::align_of;
3use std::mem::size_of;
4
5use tenferro_tensor::{
6    AllocationDomainId, DType, Placement, ShapeVec, StrideVec, Tensor, TensorRead, TensorScalar,
7    TensorView, TypedTensor, TypedTensorView,
8};
9
10use super::{InputSignatureError, LayoutClass, PrepareError};
11
12const COMPACT_COL_MAJOR_LAYOUT: &str = "tenferro.layout.compact-col-major.v1";
13const STRIDED_LAYOUT: &str = "tenferro.layout.strided.v1";
14
15/// Value-free metadata signature for a tensor input.
16///
17/// # Examples
18///
19/// ```
20/// # fn main() -> Result<(), Box<dyn std::error::Error>> {
21/// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
22/// use tenferro_tensor::Placement;
23///
24/// let entry = InputSignatureEntry::new(
25///     DType::F64,
26///     [2_usize].into_iter().collect(),
27///     Placement::default(),
28///     LayoutClass::new("tenferro.layout.strided")?,
29///     [1_isize].into_iter().collect(),
30///     Some(3),
31/// )?;
32/// assert_eq!(entry.dtype(), DType::F64);
33/// # Ok(())
34/// # }
35/// ```
36#[derive(Clone, Debug, Eq, Hash, PartialEq)]
37pub struct InputSignatureEntry {
38    dtype: DType,
39    shape: ShapeVec,
40    placement: Placement,
41    layout_class: LayoutClass,
42    strides: StrideVec,
43    alignment_log2: Option<u8>,
44    backend_family: Option<&'static str>,
45    allocation_domain: Option<AllocationDomainId>,
46}
47
48#[derive(Clone, Copy)]
49struct InputPhysicalIdentity {
50    backend_family: Option<&'static str>,
51    allocation_domain: Option<AllocationDomainId>,
52}
53
54impl InputSignatureEntry {
55    /// Build one value-free input signature entry.
56    ///
57    /// # Examples
58    ///
59    /// ```
60    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
61    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
62    /// use tenferro_tensor::Placement;
63    ///
64    /// let entry = InputSignatureEntry::new(
65    ///     DType::I32,
66    ///     [4_usize].into_iter().collect(),
67    ///     Placement::default(),
68    ///     LayoutClass::new("tenferro.layout.compact")?,
69    ///     [1_isize].into_iter().collect(),
70    ///     None,
71    /// )?;
72    /// assert_eq!(entry.shape(), &[4]);
73    /// # Ok(())
74    /// # }
75    /// ```
76    ///
77    /// # Errors
78    ///
79    /// Returns [`InputSignatureError::ShapeStrideRankMismatch`] when shape and
80    /// stride ranks differ, or [`InputSignatureError::InvalidAlignmentClass`]
81    /// when `alignment_log2` is outside the finite `usize` alignment lattice.
82    pub fn new(
83        dtype: DType,
84        shape: ShapeVec,
85        placement: Placement,
86        layout_class: LayoutClass,
87        strides: StrideVec,
88        alignment_log2: Option<u8>,
89    ) -> Result<Self, InputSignatureError> {
90        validate_entry(&shape, &strides, alignment_log2)?;
91        Ok(Self {
92            dtype,
93            shape,
94            placement,
95            layout_class,
96            strides,
97            alignment_log2,
98            backend_family: None,
99            allocation_domain: None,
100        })
101    }
102
103    fn from_validated_metadata(
104        dtype: DType,
105        shape: ShapeVec,
106        placement: Placement,
107        layout_class: LayoutClass,
108        strides: StrideVec,
109        alignment_log2: Option<u8>,
110        physical_identity: InputPhysicalIdentity,
111    ) -> Self {
112        Self {
113            dtype,
114            shape,
115            placement,
116            layout_class,
117            strides,
118            alignment_log2,
119            backend_family: physical_identity.backend_family,
120            allocation_domain: physical_identity.allocation_domain,
121        }
122    }
123
124    /// Return the dtype component.
125    ///
126    /// # Examples
127    ///
128    /// ```
129    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
130    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
131    /// use tenferro_tensor::Placement;
132    ///
133    /// let entry = InputSignatureEntry::new(
134    ///     DType::Bool,
135    ///     [1_usize].into_iter().collect(),
136    ///     Placement::default(),
137    ///     LayoutClass::new("tenferro.layout.strided")?,
138    ///     [1_isize].into_iter().collect(),
139    ///     None,
140    /// )?;
141    /// assert_eq!(entry.dtype(), DType::Bool);
142    /// # Ok(())
143    /// # }
144    /// ```
145    pub fn dtype(&self) -> DType {
146        self.dtype
147    }
148
149    /// Return the shape component.
150    ///
151    /// # Examples
152    ///
153    /// ```
154    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
155    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
156    /// use tenferro_tensor::Placement;
157    ///
158    /// let entry = InputSignatureEntry::new(
159    ///     DType::F64,
160    ///     [2_usize, 3].into_iter().collect(),
161    ///     Placement::default(),
162    ///     LayoutClass::new("tenferro.layout.strided")?,
163    ///     [1_isize, 2].into_iter().collect(),
164    ///     None,
165    /// )?;
166    /// assert_eq!(entry.shape(), &[2, 3]);
167    /// # Ok(())
168    /// # }
169    /// ```
170    pub fn shape(&self) -> &[usize] {
171        &self.shape
172    }
173
174    /// Return the placement metadata component.
175    ///
176    /// # Examples
177    ///
178    /// ```
179    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
180    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
181    /// use tenferro_tensor::{MemoryKind, Placement};
182    ///
183    /// let entry = InputSignatureEntry::new(
184    ///     DType::F64,
185    ///     [1_usize].into_iter().collect(),
186    ///     Placement::default(),
187    ///     LayoutClass::new("tenferro.layout.strided")?,
188    ///     [1_isize].into_iter().collect(),
189    ///     None,
190    /// )?;
191    /// assert_eq!(entry.placement().memory_kind, MemoryKind::UnpinnedHost);
192    /// # Ok(())
193    /// # }
194    /// ```
195    pub fn placement(&self) -> &Placement {
196        &self.placement
197    }
198
199    /// Return the layout class component.
200    ///
201    /// # Examples
202    ///
203    /// ```
204    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
205    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
206    /// use tenferro_tensor::Placement;
207    ///
208    /// let layout = LayoutClass::new("tenferro.layout.strided")?;
209    /// let entry = InputSignatureEntry::new(
210    ///     DType::F64,
211    ///     [1_usize].into_iter().collect(),
212    ///     Placement::default(),
213    ///     layout.clone(),
214    ///     [1_isize].into_iter().collect(),
215    ///     None,
216    /// )?;
217    /// assert_eq!(entry.layout_class(), &layout);
218    /// # Ok(())
219    /// # }
220    /// ```
221    pub fn layout_class(&self) -> &LayoutClass {
222        &self.layout_class
223    }
224
225    /// Return the stride metadata component.
226    ///
227    /// # Examples
228    ///
229    /// ```
230    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
231    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
232    /// use tenferro_tensor::Placement;
233    ///
234    /// let entry = InputSignatureEntry::new(
235    ///     DType::F64,
236    ///     [2_usize].into_iter().collect(),
237    ///     Placement::default(),
238    ///     LayoutClass::new("tenferro.layout.strided")?,
239    ///     [2_isize].into_iter().collect(),
240    ///     None,
241    /// )?;
242    /// assert_eq!(entry.strides(), &[2]);
243    /// # Ok(())
244    /// # }
245    /// ```
246    pub fn strides(&self) -> &[isize] {
247        &self.strides
248    }
249
250    /// Return the known alignment class, if available.
251    ///
252    /// # Examples
253    ///
254    /// ```
255    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
256    /// use tenferro_runtime::{DType, InputSignatureEntry, LayoutClass};
257    /// use tenferro_tensor::Placement;
258    ///
259    /// let entry = InputSignatureEntry::new(
260    ///     DType::F64,
261    ///     [1_usize].into_iter().collect(),
262    ///     Placement::default(),
263    ///     LayoutClass::new("tenferro.layout.strided")?,
264    ///     [1_isize].into_iter().collect(),
265    ///     Some(3),
266    /// )?;
267    /// assert_eq!(entry.alignment_log2(), Some(3));
268    /// # Ok(())
269    /// # }
270    /// ```
271    pub fn alignment_log2(&self) -> Option<u8> {
272        self.alignment_log2
273    }
274
275    pub(super) fn backend_family(&self) -> Option<&'static str> {
276        self.backend_family
277    }
278
279    pub(super) fn allocation_domain(&self) -> Option<AllocationDomainId> {
280        self.allocation_domain
281    }
282
283    pub(crate) fn logical_retained_bytes(&self) -> Option<usize> {
284        checked_sum([
285            spilled_bytes::<usize>(self.shape.spilled(), self.shape.len())?,
286            spilled_bytes::<isize>(self.strides.spilled(), self.strides.len())?,
287        ])
288    }
289}
290
291/// Value-free signature of all tensor inputs for one prepare request.
292///
293/// # Examples
294///
295/// ```
296/// use tenferro_runtime::InputSignature;
297///
298/// let signature = InputSignature::new(Vec::new());
299/// assert!(signature.entries().is_empty());
300/// ```
301#[derive(Clone, Debug, Eq, Hash, PartialEq)]
302pub struct InputSignature {
303    entries: Vec<InputSignatureEntry>,
304}
305
306impl InputSignature {
307    /// Build a signature from already prepared entries.
308    ///
309    /// # Examples
310    ///
311    /// ```
312    /// use tenferro_runtime::InputSignature;
313    ///
314    /// let signature = InputSignature::new(Vec::new());
315    /// assert_eq!(signature.entries().len(), 0);
316    /// ```
317    pub fn new(entries: Vec<InputSignatureEntry>) -> Self {
318        Self { entries }
319    }
320
321    /// Build a value-free signature from borrowed tensor reads.
322    ///
323    /// # Examples
324    ///
325    /// ```
326    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
327    /// use tenferro_runtime::{InputSignature, TensorRead, Tensor};
328    ///
329    /// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
330    /// let signature = InputSignature::from_reads(&[TensorRead::from_tensor(&tensor)])?;
331    /// assert_eq!(signature.entries()[0].shape(), &[2]);
332    /// # Ok(())
333    /// # }
334    /// ```
335    ///
336    /// # Errors
337    ///
338    /// Returns [`PrepareError::InputSignature`] with the original typed tensor
339    /// metadata error when shape, stride, or compactness metadata cannot be read.
340    pub fn from_reads(reads: &[TensorRead<'_>]) -> Result<Self, PrepareError> {
341        let mut entries = Vec::with_capacity(reads.len());
342        for (input, read) in reads.iter().enumerate() {
343            let strides = read
344                .strides()
345                .map_err(|source| PrepareError::InputSignature {
346                    source: InputSignatureError::TensorMetadata { input, source },
347                })?;
348            let compact =
349                read.is_col_major_contiguous()
350                    .map_err(|source| PrepareError::InputSignature {
351                        source: InputSignatureError::TensorMetadata { input, source },
352                    })?;
353            let shape = read.shape().iter().copied().collect();
354            entries.push(InputSignatureEntry::from_validated_metadata(
355                read.dtype(),
356                shape,
357                read_placement(read),
358                layout_class(compact),
359                strides.into_iter().collect(),
360                read_alignment_log2(read),
361                InputPhysicalIdentity {
362                    backend_family: read.backend_family(),
363                    allocation_domain: read.allocation_domain(),
364                },
365            ));
366        }
367        Ok(Self { entries })
368    }
369
370    /// Return the per-input entries.
371    ///
372    /// # Examples
373    ///
374    /// ```
375    /// use tenferro_runtime::InputSignature;
376    ///
377    /// assert!(InputSignature::new(Vec::new()).entries().is_empty());
378    /// ```
379    pub fn entries(&self) -> &[InputSignatureEntry] {
380        &self.entries
381    }
382
383    pub(crate) fn logical_retained_bytes(&self) -> Option<usize> {
384        checked_sum([
385            self.entries
386                .len()
387                .checked_mul(size_of::<InputSignatureEntry>())?,
388            checked_sum_options(
389                self.entries
390                    .iter()
391                    .map(InputSignatureEntry::logical_retained_bytes),
392            )?,
393        ])
394    }
395}
396
397fn spilled_bytes<T>(spilled: bool, len: usize) -> Option<usize> {
398    if spilled {
399        len.checked_mul(size_of::<T>())
400    } else {
401        Some(0)
402    }
403}
404
405fn checked_sum(values: impl IntoIterator<Item = usize>) -> Option<usize> {
406    values
407        .into_iter()
408        .try_fold(0usize, |sum, value| sum.checked_add(value))
409}
410
411fn checked_sum_options(values: impl IntoIterator<Item = Option<usize>>) -> Option<usize> {
412    values
413        .into_iter()
414        .try_fold(0usize, |sum, value| sum.checked_add(value?))
415}
416
417fn validate_entry(
418    shape: &[usize],
419    strides: &[isize],
420    alignment_log2: Option<u8>,
421) -> Result<(), InputSignatureError> {
422    if shape.len() != strides.len() {
423        return Err(InputSignatureError::ShapeStrideRankMismatch {
424            rank: shape.len(),
425            stride_count: strides.len(),
426        });
427    }
428    if let Some(alignment_log2) = alignment_log2
429        && u32::from(alignment_log2) >= usize::BITS
430    {
431        return Err(InputSignatureError::InvalidAlignmentClass { alignment_log2 });
432    }
433    Ok(())
434}
435
436pub(super) fn read_placement(read: &TensorRead<'_>) -> Placement {
437    match read {
438        TensorRead::Tensor(tensor) => tensor.placement().clone(),
439        TensorRead::View(view) => view_placement(view),
440    }
441}
442
443fn view_placement(view: &TensorView<'_>) -> Placement {
444    match view {
445        TensorView::F32(view) => view.placement().clone(),
446        TensorView::F64(view) => view.placement().clone(),
447        TensorView::I32(view) => view.placement().clone(),
448        TensorView::I64(view) => view.placement().clone(),
449        TensorView::Bool(view) => view.placement().clone(),
450        TensorView::C32(view) => view.placement().clone(),
451        TensorView::C64(view) => view.placement().clone(),
452    }
453}
454
455fn layout_class(compact: bool) -> LayoutClass {
456    let value = if compact {
457        COMPACT_COL_MAJOR_LAYOUT
458    } else {
459        STRIDED_LAYOUT
460    };
461    LayoutClass::runtime_created(value)
462}
463
464fn read_alignment_log2(read: &TensorRead<'_>) -> Option<u8> {
465    match read {
466        TensorRead::Tensor(tensor) => tensor_alignment_log2(tensor),
467        TensorRead::View(view) => view_alignment_log2(view),
468    }
469}
470
471fn tensor_alignment_log2(tensor: &Tensor) -> Option<u8> {
472    match tensor.dtype() {
473        DType::F32 => tensor
474            .as_typed::<f32>()
475            .and_then(typed_tensor_alignment_log2),
476        DType::F64 => tensor
477            .as_typed::<f64>()
478            .and_then(typed_tensor_alignment_log2),
479        DType::I32 => tensor
480            .as_typed::<i32>()
481            .and_then(typed_tensor_alignment_log2),
482        DType::I64 => tensor
483            .as_typed::<i64>()
484            .and_then(typed_tensor_alignment_log2),
485        DType::Bool => tensor
486            .as_typed::<bool>()
487            .and_then(typed_tensor_alignment_log2),
488        DType::C32 => tensor
489            .as_typed::<Complex32>()
490            .and_then(typed_tensor_alignment_log2),
491        DType::C64 => tensor
492            .as_typed::<Complex64>()
493            .and_then(typed_tensor_alignment_log2),
494        // A caller-owned payload's owner declares its own alignment, which is also
495        // what the accessor falls back to when the tag and the runtime dtype differ.
496        DType::External(_) => None,
497    }
498}
499
500fn typed_tensor_alignment_log2<T: TensorScalar>(tensor: &TypedTensor<T>) -> Option<u8> {
501    if tensor.backend_family().is_some() {
502        return None;
503    }
504    if shape_is_empty(tensor.shape()) {
505        return Some(type_alignment_log2::<T>());
506    }
507    tensor
508        .host_data()
509        .ok()
510        .map(|data| pointer_alignment_log2::<T>(data.as_ptr()))
511}
512
513fn view_alignment_log2(view: &TensorView<'_>) -> Option<u8> {
514    match view {
515        TensorView::F32(view) => typed_view_alignment_log2(view),
516        TensorView::F64(view) => typed_view_alignment_log2(view),
517        TensorView::I32(view) => typed_view_alignment_log2(view),
518        TensorView::I64(view) => typed_view_alignment_log2(view),
519        TensorView::Bool(view) => typed_view_alignment_log2(view),
520        TensorView::C32(view) => typed_view_alignment_log2(view),
521        TensorView::C64(view) => typed_view_alignment_log2(view),
522    }
523}
524
525fn typed_view_alignment_log2<T: TensorScalar + 'static>(
526    view: &TypedTensorView<'_, T>,
527) -> Option<u8> {
528    if view.backend_family().is_some() {
529        return None;
530    }
531    if shape_is_empty(view.shape()) {
532        return Some(type_alignment_log2::<T>());
533    }
534    view.host_storage().ok().map(|data| {
535        let pointer = data.as_ptr().wrapping_offset(view.offset());
536        pointer_alignment_log2::<T>(pointer)
537    })
538}
539
540fn shape_is_empty(shape: &[usize]) -> bool {
541    shape.contains(&0)
542}
543
544fn type_alignment_log2<T>() -> u8 {
545    align_of::<T>().trailing_zeros().min(usize::BITS - 1) as u8
546}
547
548fn pointer_alignment_log2<T>(pointer: *const T) -> u8 {
549    (pointer as usize)
550        .trailing_zeros()
551        .min(align_of::<T>().trailing_zeros())
552        .min(usize::BITS - 1) as u8
553}