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}