1use std::fmt;
2use std::marker::PhantomData;
3use std::mem::{align_of, size_of};
4use std::ptr::NonNull;
5
6use crate::types::{tensor_from_group, tensor_view_from_group};
7use crate::{DType, DynRank, Placement, TensorLayout, TensorRank, TensorRead, TensorScalar};
8use smallvec::SmallVec;
9
10use super::identity::RootResourceId;
11use super::prepared::{
12 prepare_read, prepare_write, validate_descriptor, AccessError, AccessTarget, CheckedDescriptor,
13 CheckedRead, CheckedWrite, PreparedRead, PreparedWrite, ProviderReadMapping,
14 ProviderWriteMapping, WriteInjectivityProof,
15};
16use super::root::{BackendAllocation, OwnedStorage, ProviderKind};
17use super::span::{ByteRange, RootBoundSpan};
18
19#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
21pub(crate) struct AllocationSlot(u32);
22
23impl AllocationSlot {
24 pub(crate) const fn index(self) -> usize {
25 self.0 as usize
26 }
27
28 #[cfg(test)]
29 pub(crate) const fn test_raw(raw: u32) -> Self {
30 Self(raw)
31 }
32}
33
34#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
45pub struct DescriptorSlot(u32);
46
47impl DescriptorSlot {
48 pub const fn index(self) -> usize {
59 self.0 as usize
60 }
61
62 pub fn from_index(index: usize) -> Option<Self> {
73 match u32::try_from(index) {
74 Ok(index) => Some(Self(index)),
75 Err(_) => None,
76 }
77 }
78
79 #[cfg(test)]
80 pub(crate) const fn test_raw(raw: u32) -> Self {
81 Self(raw)
82 }
83}
84
85#[derive(Clone, Debug, PartialEq, Eq)]
87pub(crate) struct DescriptorInput<R: TensorRank> {
88 relative: ByteRange,
89 shape: R::Shape,
90 strides: R::Strides,
91 offset: isize,
92 require_injective: bool,
93}
94
95impl<R: TensorRank> DescriptorInput<R> {
96 pub(crate) fn new(
97 relative: ByteRange,
98 shape: R::Shape,
99 strides: R::Strides,
100 offset: isize,
101 require_injective: bool,
102 ) -> Self {
103 Self {
104 relative,
105 shape,
106 strides,
107 offset,
108 require_injective,
109 }
110 }
111}
112
113#[derive(Clone, Debug, PartialEq, Eq)]
119pub(crate) struct DescriptorRecord {
120 allocation: AllocationSlot,
121 element_count: usize,
122 provider: ProviderKind,
123 placement: Placement,
124 envelope: Option<ByteRange>,
125 write_injective: bool,
126 checked: CheckedDescriptor<DynRank>,
127}
128
129impl DescriptorRecord {
130 pub(crate) const fn allocation(&self) -> AllocationSlot {
131 self.allocation
132 }
133
134 pub(crate) const fn span(&self) -> RootBoundSpan {
135 self.checked.span()
136 }
137
138 pub(crate) const fn dtype(&self) -> DType {
139 self.checked.dtype()
140 }
141
142 pub(crate) const fn element_size(&self) -> usize {
143 self.checked.element_size()
144 }
145
146 pub(crate) const fn element_count(&self) -> usize {
147 self.element_count
148 }
149
150 pub(crate) fn layout(&self) -> &TensorLayout<DynRank> {
151 self.checked.logical_layout()
152 }
153
154 pub(crate) const fn provider(&self) -> ProviderKind {
155 self.provider
156 }
157
158 pub(crate) fn placement(&self) -> &Placement {
159 &self.placement
160 }
161
162 pub(crate) const fn envelope(&self) -> Option<ByteRange> {
163 self.envelope
164 }
165
166 pub(crate) const fn write_injective(&self) -> bool {
167 self.write_injective
168 }
169}
170
171fn into_group_parts(
179 tensor: crate::Tensor,
180) -> Result<(AllocationGroup, DescriptorSlot), GroupError> {
181 if tensor.external_payload().is_some() {
182 return Err(GroupError::InvalidDescriptor {
183 message: "a caller-owned payload has no allocation group".to_owned(),
184 });
185 }
186 Ok(tensor.into_group_parts())
187}
188
189#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
200pub enum GroupError {
201 #[error("group index overflows u32")]
202 IndexOverflow,
203 #[error("descriptor slot {slot} is outside the group")]
204 DescriptorSlotOutOfBounds { slot: usize },
205 #[error("descriptor slot {slot} is vacant")]
206 DescriptorSlotVacant { slot: usize },
207 #[error("allocation slot {slot} is outside the group")]
208 AllocationSlotOutOfBounds { slot: usize },
209 #[error("allocation slot {slot} is vacant")]
210 AllocationSlotVacant { slot: usize },
211 #[error("descriptor validation failed: {message}")]
212 InvalidDescriptor { message: String },
213 #[error("descriptor dtype mismatch: expected {expected:?}, actual {actual:?}")]
214 DTypeMismatch { expected: DType, actual: DType },
215 #[error("descriptor rank mismatch: expected {expected}, actual {actual}")]
216 RankMismatch { expected: usize, actual: usize },
217 #[error("allocation slot {allocation} has more than one live descriptor")]
218 AliasedAllocation { allocation: usize },
219 #[error(
223 "descriptor slot {slot} is not a compact whole-allocation layout; materialize it instead of extracting"
224 )]
225 NonCompactDescriptor { slot: usize },
226}
227
228#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
230pub(crate) enum DisjointViewError {
231 #[error(transparent)]
232 Group(#[from] GroupError),
233 #[error("descriptor slot {slot} appears more than once")]
234 DuplicateSlot { slot: usize },
235 #[error("descriptor slot {slot} has a non-injective mutable layout")]
236 NonInjective { slot: usize },
237 #[error("requested mutable descriptor envelopes overlap")]
238 PairwiseOverlap,
239 #[error("requested mutable descriptors are not provably disjoint")]
240 NotProvablyDisjoint,
241}
242
243#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
245pub(crate) enum ExtractError {
246 #[error(transparent)]
247 Group(#[from] GroupError),
248 #[error("allocation slot {allocation} still has another descriptor")]
249 AliasedAllocation { allocation: usize },
250}
251
252#[derive(Default)]
265pub struct AllocationGroup {
266 allocations: SmallVec<[Option<OwnedStorage>; 1]>,
271 descriptors: SmallVec<[Option<DescriptorRecord>; 1]>,
272}
273
274pub(crate) struct GroupReadView<'a, T, R: TensorRank> {
276 owner: NonNull<OwnedStorage>,
277 descriptor: DescriptorRecord,
278 _borrow: PhantomData<(&'a OwnedStorage, T, R)>,
279}
280
281unsafe impl<'a, T: Send, R: TensorRank> Send for GroupReadView<'a, T, R> {}
284unsafe impl<'a, T: Sync, R: TensorRank> Sync for GroupReadView<'a, T, R> {}
285
286impl<'a, T, R: TensorRank> Clone for GroupReadView<'a, T, R> {
287 fn clone(&self) -> Self {
288 Self {
289 owner: self.owner,
290 descriptor: self.descriptor.clone(),
291 _borrow: PhantomData,
292 }
293 }
294}
295
296impl<'a, T, R: TensorRank> GroupReadView<'a, T, R> {
297 pub(crate) fn clone_dyn(&self) -> GroupReadView<'a, T, crate::DynRank> {
298 GroupReadView {
299 owner: self.owner,
300 descriptor: self.descriptor.clone(),
301 _borrow: PhantomData,
302 }
303 }
304
305 pub(crate) fn into_dyn(self) -> GroupReadView<'a, T, crate::DynRank> {
306 GroupReadView {
307 owner: self.owner,
308 descriptor: self.descriptor,
309 _borrow: PhantomData,
310 }
311 }
312}
313
314impl<T, R: TensorRank> std::fmt::Debug for GroupReadView<'_, T, R> {
315 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
316 formatter
317 .debug_struct("GroupReadView")
318 .field("descriptor", &self.descriptor)
319 .finish_non_exhaustive()
320 }
321}
322
323impl<'a, T: TensorScalar, R: TensorRank> GroupReadView<'a, T, R> {
324 pub(crate) fn descriptor(&self) -> &DescriptorRecord {
325 &self.descriptor
326 }
327
328 pub(crate) fn provider_kind(&self) -> ProviderKind {
329 self.descriptor.provider
330 }
331
332 pub(crate) fn backend_identity(
333 &self,
334 ) -> Option<(crate::AllocationDomainId, crate::AllocationId)> {
335 if self.descriptor.provider == ProviderKind::Cpu {
336 return None;
337 }
338 let key = unsafe { self.owner.as_ref().root_identity().extent().key() };
339 Some((key.domain(), key.local()))
340 }
341
342 pub(crate) fn storage_buffer(&self) -> Option<&'a crate::StorageBuffer<T>> {
343 let buffer = unsafe {
346 self.owner
347 .as_ref()
348 .host_buffer::<T>()
349 .or_else(|| self.owner.as_ref().backend_buffer::<T>())
350 }?;
351 Some(unsafe {
352 std::mem::transmute::<&crate::StorageBuffer<T>, &'a crate::StorageBuffer<T>>(buffer)
353 })
354 }
355
356 pub(crate) fn map_read(&self) -> Result<ProviderReadMapping<'_>, AccessError> {
357 unsafe {
360 self.owner
361 .as_ref()
362 .as_ref()
363 .map_read(self.descriptor.span(), self.descriptor.dtype())
364 }
365 }
366
367 pub(crate) fn prepare_device_read_for_layout(
368 &self,
369 layout: &TensorLayout<R>,
370 ) -> Result<Box<dyn crate::PreparedDeviceAccess + 'a>, AccessError> {
371 self.prepare_read_for_layout(layout, AccessTarget::Device)?
372 .into_device_state()
373 }
374
375 pub(crate) fn prepare_host_read_for_layout(
376 &self,
377 layout: &TensorLayout<R>,
378 ) -> Result<PreparedRead<'a, T, R>, AccessError> {
379 self.prepare_read_for_layout(layout, AccessTarget::Host)
380 }
381
382 fn prepare_read_for_layout(
383 &self,
384 layout: &TensorLayout<R>,
385 target: AccessTarget,
386 ) -> Result<PreparedRead<'a, T, R>, AccessError> {
387 let owner: crate::storage::root::StorageRef<'a> = unsafe { self.owner.as_ref().as_ref() };
389 let checked: CheckedRead<'a, R> = CheckedRead::new::<T>(
390 owner,
392 self.descriptor.span(),
393 R::shape_from_vec(layout.shape().iter().copied().collect()).map_err(|error| {
394 AccessError::InvalidLayout {
395 message: error.to_string(),
396 }
397 })?,
398 R::strides_from_vec(layout.strides().iter().copied().collect()).map_err(|error| {
399 AccessError::InvalidLayout {
400 message: error.to_string(),
401 }
402 })?,
403 layout.offset(),
404 )?;
405 prepare_read::<T, R>(checked, target).map_err(|failure| failure.1)
406 }
407
408 pub(crate) fn prepare_host_read(&self) -> Result<PreparedRead<'_, T, DynRank>, AccessError> {
409 prepare_read(
410 CheckedRead::from_validated(
411 unsafe { self.owner.as_ref().as_ref() },
413 self.descriptor.checked.clone(),
414 ),
415 AccessTarget::Host,
416 )
417 .map_err(|failure| failure.1)
418 }
419
420 pub(crate) fn host_slice(&self) -> Result<&'a [T], AccessError> {
421 unsafe {
423 self.owner
424 .as_ref()
425 .as_ref()
426 .host_slice(self.descriptor.span(), self.descriptor.dtype())
427 }
428 }
429}
430
431impl<'a, T: 'static, R: TensorRank> GroupReadView<'a, T, R> {
432 pub(crate) fn backend_allocation(&self) -> Option<&'a dyn BackendAllocation> {
433 let allocation = unsafe { self.owner.as_ref().backend_allocation() }?;
435 Some(unsafe {
436 std::mem::transmute::<&dyn BackendAllocation, &'a dyn BackendAllocation>(allocation)
437 })
438 }
439
440 pub(crate) fn backend_buffer(&self) -> Option<&'a crate::StorageBuffer<T>> {
441 let buffer = unsafe { self.owner.as_ref().backend_buffer::<T>() }?;
444 Some(unsafe {
445 std::mem::transmute::<&crate::StorageBuffer<T>, &'a crate::StorageBuffer<T>>(buffer)
446 })
447 }
448}
449
450pub(crate) struct GroupWriteView<'a, T, R: TensorRank> {
454 owner: NonNull<OwnedStorage>,
455 descriptor: DescriptorRecord,
456 _borrow: PhantomData<(&'a mut [u8], T, R)>,
457}
458
459unsafe impl<'a, T: Send, R: TensorRank> Send for GroupWriteView<'a, T, R> {}
462
463impl<T, R: TensorRank> std::fmt::Debug for GroupWriteView<'_, T, R> {
464 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
465 formatter
466 .debug_struct("GroupWriteView")
467 .field("descriptor", &self.descriptor)
468 .finish_non_exhaustive()
469 }
470}
471
472impl<'a, T: TensorScalar, R: TensorRank> GroupWriteView<'a, T, R> {
473 pub(crate) fn descriptor(&self) -> &DescriptorRecord {
474 &self.descriptor
475 }
476
477 pub(crate) fn map_write(&mut self) -> Result<ProviderWriteMapping<'_>, AccessError> {
478 unsafe {
482 self.owner
483 .as_mut()
484 .as_mut()
485 .map_write(self.descriptor.span(), self.descriptor.dtype())
486 }
487 }
488
489 pub(crate) fn backend_buffer_mut(&mut self) -> Option<&'a mut crate::StorageBuffer<T>> {
490 let owner = unsafe { &mut *self.owner.as_ptr() };
493 let buffer = owner.backend_buffer_mut::<T>()?;
494 Some(unsafe {
495 std::mem::transmute::<&mut crate::StorageBuffer<T>, &'a mut crate::StorageBuffer<T>>(
496 buffer,
497 )
498 })
499 }
500
501 pub(crate) fn prepare_device_write_for_layout(
502 &mut self,
503 layout: &TensorLayout<R>,
504 ) -> Result<Box<dyn crate::PreparedDeviceAccess + 'a>, AccessError> {
505 let owner: crate::storage::root::StorageMut<'a> =
506 unsafe { (&mut *self.owner.as_ptr()).as_mut() };
507 let checked: CheckedWrite<'a, R> = CheckedWrite::new::<T>(
508 owner,
510 self.descriptor.span(),
511 R::shape_from_vec(layout.shape().iter().copied().collect()).map_err(|error| {
512 AccessError::InvalidLayout {
513 message: error.to_string(),
514 }
515 })?,
516 R::strides_from_vec(layout.strides().iter().copied().collect()).map_err(|error| {
517 AccessError::InvalidLayout {
518 message: error.to_string(),
519 }
520 })?,
521 layout.offset(),
522 )?;
523 prepare_write::<T, R>(checked, AccessTarget::Device)
524 .map_err(|failure| failure.1)?
525 .into_device_state()
526 }
527
528 pub(crate) fn prepare_host_write(
529 &mut self,
530 ) -> Result<PreparedWrite<'_, T, DynRank>, AccessError> {
531 let checked = CheckedWrite::from_validated(
532 unsafe { self.owner.as_mut().as_mut() },
534 self.descriptor.checked.clone(),
535 WriteInjectivityProof,
536 );
537 prepare_write(checked, AccessTarget::Host).map_err(|failure| failure.1)
538 }
539
540 pub(crate) fn host_slice_mut(&mut self) -> Result<&'a mut [T], AccessError> {
541 unsafe {
544 self.owner
545 .as_mut()
546 .as_mut()
547 .host_slice_mut(self.descriptor.span(), self.descriptor.dtype())
548 }
549 }
550}
551
552impl<'a, T: 'static, R: TensorRank> GroupWriteView<'a, T, R> {
553 pub(crate) fn backend_buffer(&self) -> Option<&crate::StorageBuffer<T>> {
554 unsafe { self.owner.as_ref().backend_buffer::<T>() }
557 }
558}
559
560impl fmt::Debug for AllocationGroup {
561 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
562 formatter
563 .debug_struct("AllocationGroup")
564 .field("allocation_count", &self.allocations.len())
565 .field("descriptor_count", &self.descriptors.len())
566 .finish()
567 }
568}
569
570impl AllocationGroup {
571 pub(crate) fn new() -> Self {
572 Self::default()
573 }
574
575 pub fn from_tensors(
594 tensors: Vec<crate::Tensor>,
595 ) -> Result<(Self, Box<[DescriptorSlot]>), GroupError> {
596 let mut group = Self::new();
597 let mut bindings = Vec::with_capacity(tensors.len());
598 for tensor in tensors {
599 let (source, source_slot) = into_group_parts(tensor)?;
600 bindings.push(group.append_group(source, source_slot)?);
601 }
602 Ok((group, bindings.into_boxed_slice()))
603 }
604
605 pub fn read_views<'a>(
614 &'a self,
615 bindings: &[DescriptorSlot],
616 ) -> Result<Vec<TensorRead<'a>>, GroupError> {
617 bindings
618 .iter()
619 .map(|&slot| self.tensor_read(slot))
620 .collect()
621 }
622
623 pub fn read_view<'a>(&'a self, slot: DescriptorSlot) -> Result<TensorRead<'a>, GroupError> {
632 self.tensor_read(slot)
633 }
634
635 fn tensor_read<'a>(&'a self, slot: DescriptorSlot) -> Result<TensorRead<'a>, GroupError> {
636 let dtype = self.resolve_descriptor(slot)?.1.dtype();
637 let view = match dtype {
638 DType::F32 => tensor_view_from_group(self.view::<f32, DynRank>(slot)?),
639 DType::F64 => tensor_view_from_group(self.view::<f64, DynRank>(slot)?),
640 DType::I32 => tensor_view_from_group(self.view::<i32, DynRank>(slot)?),
641 DType::I64 => tensor_view_from_group(self.view::<i64, DynRank>(slot)?),
642 DType::Bool => tensor_view_from_group(self.view::<bool, DynRank>(slot)?),
643 DType::C32 => {
644 tensor_view_from_group(self.view::<num_complex::Complex32, DynRank>(slot)?)
645 }
646 DType::C64 => {
647 tensor_view_from_group(self.view::<num_complex::Complex64, DynRank>(slot)?)
648 }
649 DType::External(_) => unreachable!("descriptors are preset-typed"),
652 }
653 .map_err(|error| GroupError::InvalidDescriptor {
654 message: error.to_string(),
655 })?;
656 Ok(TensorRead::from_view(view))
657 }
658
659 pub fn append_tensor(&mut self, tensor: crate::Tensor) -> Result<DescriptorSlot, GroupError> {
668 let (source, source_slot) = into_group_parts(tensor)?;
669 self.append_group(source, source_slot)
670 }
671
672 pub fn append_group(
679 &mut self,
680 mut source: AllocationGroup,
681 source_slot: DescriptorSlot,
682 ) -> Result<DescriptorSlot, GroupError> {
683 let allocation_offset =
684 u32::try_from(self.allocations.len()).map_err(|_| GroupError::IndexOverflow)?;
685 let descriptor_offset =
686 u32::try_from(self.descriptors.len()).map_err(|_| GroupError::IndexOverflow)?;
687 source.resolve_descriptor(source_slot)?;
688
689 for owner in source.allocations.drain(..) {
690 self.allocations.push(owner);
691 }
692 for descriptor in source.descriptors.drain(..) {
693 let descriptor = match descriptor {
694 Some(mut descriptor) => {
695 let allocation = descriptor
696 .allocation
697 .0
698 .checked_add(allocation_offset)
699 .ok_or(GroupError::IndexOverflow)?;
700 descriptor.allocation = AllocationSlot(allocation);
701 Some(descriptor)
702 }
703 None => None,
704 };
705 self.descriptors.push(descriptor);
706 }
707
708 let source_descriptor_index =
709 u32::try_from(source_slot.index()).map_err(|_| GroupError::IndexOverflow)?;
710 Ok(DescriptorSlot(
711 source_descriptor_index
712 .checked_add(descriptor_offset)
713 .ok_or(GroupError::IndexOverflow)?,
714 ))
715 }
716
717 pub(crate) fn set_descriptor_placement(
718 &mut self,
719 slot: DescriptorSlot,
720 placement: Placement,
721 ) -> Result<(), GroupError> {
722 let descriptor = self
723 .descriptors
724 .get_mut(slot.index())
725 .ok_or(GroupError::DescriptorSlotOutOfBounds { slot: slot.index() })?
726 .as_mut()
727 .ok_or(GroupError::DescriptorSlotVacant { slot: slot.index() })?;
728 descriptor.placement = placement;
729 Ok(())
730 }
731
732 pub(crate) fn publish_live_descriptor_placement(
741 &mut self,
742 slot: DescriptorSlot,
743 placement: Placement,
744 ) {
745 let descriptor = self
746 .descriptors
747 .get_mut(slot.index())
748 .and_then(Option::as_mut)
749 .unwrap_or_else(|| unreachable!("an owned group always carries its descriptor"));
750 descriptor.placement = placement;
751 }
752
753 pub(crate) fn from_host_vec<T: TensorScalar, R: TensorRank>(
754 shape: R::Shape,
755 data: Vec<T>,
756 ) -> Result<(Self, DescriptorSlot), GroupError> {
757 let owner =
758 super::root::import_host_vec(data).map_err(|error| GroupError::InvalidDescriptor {
759 message: error.to_string(),
760 })?;
761 let span = owner.root_span();
762 let mut group = Self::new();
763 let allocation = group.insert_owner(owner)?;
764 let layout =
765 TensorLayout::<R>::compact(shape).map_err(|error| GroupError::InvalidDescriptor {
766 message: error.to_string(),
767 })?;
768 let input = DescriptorInput::new(
769 ByteRange::new(0, span.byte_len()),
770 R::shape_from_vec(layout.shape().iter().copied().collect()).map_err(|error| {
771 GroupError::InvalidDescriptor {
772 message: error.to_string(),
773 }
774 })?,
775 R::strides_from_vec(layout.strides().iter().copied().collect()).map_err(|error| {
776 GroupError::InvalidDescriptor {
777 message: error.to_string(),
778 }
779 })?,
780 layout.offset(),
781 true,
782 );
783 let slot = group.insert_descriptor::<T, R>(allocation, input)?;
784 Ok((group, slot))
785 }
786
787 #[doc(hidden)]
789 pub fn from_backend_allocation<T: TensorScalar, R: TensorRank>(
790 shape: R::Shape,
791 allocation: Box<dyn BackendAllocation>,
792 ) -> Result<(Self, DescriptorSlot), GroupError> {
793 let owner = super::root::import_unique_root(allocation).map_err(|error| {
794 GroupError::InvalidDescriptor {
795 message: error.to_string(),
796 }
797 })?;
798 let span = owner.root_span();
799 let mut group = Self::new();
800 let allocation = group.insert_owner(owner)?;
801 let layout =
802 TensorLayout::<R>::compact(shape).map_err(|error| GroupError::InvalidDescriptor {
803 message: error.to_string(),
804 })?;
805 let input = DescriptorInput::new(
806 ByteRange::new(0, span.byte_len()),
807 R::shape_from_vec(layout.shape().iter().copied().collect()).map_err(|error| {
808 GroupError::InvalidDescriptor {
809 message: error.to_string(),
810 }
811 })?,
812 R::strides_from_vec(layout.strides().iter().copied().collect()).map_err(|error| {
813 GroupError::InvalidDescriptor {
814 message: error.to_string(),
815 }
816 })?,
817 layout.offset(),
818 true,
819 );
820 let slot = group.insert_descriptor::<T, R>(allocation, input)?;
821 Ok((group, slot))
822 }
823
824 pub(crate) fn from_backend_buffer<T: TensorScalar, R: TensorRank>(
825 shape: R::Shape,
826 buffer: crate::StorageBuffer<T>,
827 ) -> Result<(Self, DescriptorSlot), GroupError> {
828 let owner = super::root::import_backend_buffer(buffer).map_err(|error| {
829 GroupError::InvalidDescriptor {
830 message: error.to_string(),
831 }
832 })?;
833 let span = owner.root_span();
834 let mut group = Self::new();
835 let allocation = group.insert_owner(owner)?;
836 let layout =
837 TensorLayout::<R>::compact(shape).map_err(|error| GroupError::InvalidDescriptor {
838 message: error.to_string(),
839 })?;
840 let input = DescriptorInput::new(
841 ByteRange::new(0, span.byte_len()),
842 R::shape_from_vec(layout.shape().iter().copied().collect()).map_err(|error| {
843 GroupError::InvalidDescriptor {
844 message: error.to_string(),
845 }
846 })?,
847 R::strides_from_vec(layout.strides().iter().copied().collect()).map_err(|error| {
848 GroupError::InvalidDescriptor {
849 message: error.to_string(),
850 }
851 })?,
852 layout.offset(),
853 true,
854 );
855 let slot = group.insert_descriptor::<T, R>(allocation, input)?;
856 Ok((group, slot))
857 }
858
859 pub(crate) fn from_backend_root<T: Send + Sync + 'static>(
860 buffer: crate::StorageBuffer<T>,
861 ) -> Result<Self, GroupError> {
862 let owner = super::root::import_backend_buffer(buffer).map_err(|error| {
863 GroupError::InvalidDescriptor {
864 message: error.to_string(),
865 }
866 })?;
867 let mut group = Self::new();
868 group.insert_owner(owner)?;
869 Ok(group)
870 }
871
872 pub(crate) fn insert_owner(
873 &mut self,
874 owner: OwnedStorage,
875 ) -> Result<AllocationSlot, GroupError> {
876 let index = self.allocations.len();
877 let slot = u32::try_from(index).map_err(|_| GroupError::IndexOverflow)?;
878 self.allocations.push(Some(owner));
879 Ok(AllocationSlot(slot))
880 }
881
882 pub(crate) fn insert_descriptor<T: TensorScalar, R: TensorRank>(
883 &mut self,
884 allocation: AllocationSlot,
885 input: DescriptorInput<R>,
886 ) -> Result<DescriptorSlot, GroupError> {
887 let allocation_index = allocation.index();
888 let owner = self
889 .allocations
890 .get(allocation_index)
891 .ok_or(GroupError::AllocationSlotOutOfBounds {
892 slot: allocation_index,
893 })?
894 .as_ref()
895 .ok_or(GroupError::AllocationSlotVacant {
896 slot: allocation_index,
897 })?;
898
899 let root = owner.as_ref().root_identity();
900 let span = root.bind_relative_range(input.relative).map_err(|error| {
901 GroupError::InvalidDescriptor {
902 message: error.to_string(),
903 }
904 })?;
905 let element_size = size_of::<T>();
906 if element_size == 0 || !span.byte_len().is_multiple_of(element_size) {
907 return Err(GroupError::InvalidDescriptor {
908 message: format!(
909 "byte span {} is not divisible by element size {}",
910 span.byte_len(),
911 element_size
912 ),
913 });
914 }
915 if !span
916 .guaranteed_alignment()
917 .get()
918 .is_multiple_of(align_of::<T>())
919 {
920 return Err(GroupError::InvalidDescriptor {
921 message: format!(
922 "span alignment {} is insufficient for {}-byte alignment",
923 span.guaranteed_alignment().get(),
924 align_of::<T>()
925 ),
926 });
927 }
928
929 let shape = R::shape_into_vec(input.shape);
930 let strides = R::strides_into_vec(input.strides);
931 let layout = TensorLayout::<DynRank>::from_parts(
932 shape,
933 strides,
934 input.offset,
935 span.byte_len() / element_size,
936 )
937 .map_err(|error| GroupError::InvalidDescriptor {
938 message: error.to_string(),
939 })?;
940 let element_count = logical_element_count(layout.shape())?;
941 let write_injective = if input.require_injective {
942 layout.validate_mutable_no_overlap().map_err(|error| {
943 GroupError::InvalidDescriptor {
944 message: error.to_string(),
945 }
946 })?;
947 true
948 } else {
949 false
950 };
951 let envelope = reachable_envelope(&span, &layout, element_size)?;
952 let (checked, _) = validate_descriptor::<T, DynRank>(
953 root.root_resource(),
954 span,
955 layout.shape().iter().copied().collect(),
956 layout.strides().iter().copied().collect(),
957 layout.offset(),
958 false,
959 )
960 .map_err(|error| GroupError::InvalidDescriptor {
961 message: error.to_string(),
962 })?;
963 let record = DescriptorRecord {
964 allocation,
965 element_count,
966 provider: owner.as_ref().provider_kind(),
967 placement: Placement::default(),
968 envelope,
969 write_injective,
970 checked,
971 };
972
973 let descriptor_index = self.descriptors.len();
974 let slot = u32::try_from(descriptor_index).map_err(|_| GroupError::IndexOverflow)?;
975 self.descriptors.push(Some(record));
976 Ok(DescriptorSlot(slot))
977 }
978
979 #[allow(clippy::result_large_err)]
984 pub(crate) fn update_descriptor_layout(
985 mut self,
986 slot: DescriptorSlot,
987 shape: Vec<usize>,
988 strides: Vec<isize>,
989 offset: isize,
990 ) -> Result<Self, (Self, GroupError)> {
991 let dtype = match self.resolve_descriptor(slot) {
992 Ok((_, descriptor)) => descriptor.dtype(),
993 Err(error) => return Err((self, error)),
994 };
995 if let Some(Some(descriptor)) = self.descriptors.get_mut(slot.index()) {
996 descriptor.write_injective = false;
999 }
1000 match dtype {
1001 DType::F32 => self.reinterpret_descriptor::<f32, f32>(slot, shape, strides, offset),
1002 DType::F64 => self.reinterpret_descriptor::<f64, f64>(slot, shape, strides, offset),
1003 DType::I32 => self.reinterpret_descriptor::<i32, i32>(slot, shape, strides, offset),
1004 DType::I64 => self.reinterpret_descriptor::<i64, i64>(slot, shape, strides, offset),
1005 DType::Bool => self.reinterpret_descriptor::<bool, bool>(slot, shape, strides, offset),
1006 DType::C32 => self
1007 .reinterpret_descriptor::<num_complex::Complex32, num_complex::Complex32>(
1008 slot, shape, strides, offset,
1009 ),
1010 DType::C64 => self
1011 .reinterpret_descriptor::<num_complex::Complex64, num_complex::Complex64>(
1012 slot, shape, strides, offset,
1013 ),
1014 DType::External(_) => unreachable!("descriptors are preset-typed"),
1017 }
1018 }
1019
1020 #[allow(clippy::result_large_err)]
1028 pub(crate) fn reinterpret_descriptor<T: TensorScalar, U: TensorScalar>(
1029 mut self,
1030 slot: DescriptorSlot,
1031 shape: Vec<usize>,
1032 strides: Vec<isize>,
1033 offset: isize,
1034 ) -> Result<Self, (Self, GroupError)> {
1035 let (descriptor_index, descriptor) = match self.resolve_descriptor(slot) {
1036 Ok((index, descriptor)) => (index, descriptor.clone()),
1037 Err(error) => return Err((self, error)),
1038 };
1039 if descriptor.dtype() != T::dtype() {
1040 return Err((
1041 self,
1042 GroupError::DTypeMismatch {
1043 expected: descriptor.dtype(),
1044 actual: T::dtype(),
1045 },
1046 ));
1047 }
1048 let allocation = descriptor.allocation;
1049 let references = self
1050 .descriptors
1051 .iter()
1052 .flatten()
1053 .filter(|candidate| candidate.allocation == allocation)
1054 .count();
1055 if references != 1 {
1056 return Err((
1057 self,
1058 GroupError::AliasedAllocation {
1059 allocation: allocation.index(),
1060 },
1061 ));
1062 }
1063
1064 let root = match self.allocation_root_resource(allocation) {
1065 Some(root) => root,
1066 None => {
1067 return Err((
1068 self,
1069 GroupError::AllocationSlotOutOfBounds {
1070 slot: allocation.index(),
1071 },
1072 ))
1073 }
1074 };
1075 let span = descriptor.span();
1076 let target_element_size = std::mem::size_of::<U>();
1077 if target_element_size == 0 || !span.byte_len().is_multiple_of(target_element_size) {
1078 return Err((
1079 self,
1080 GroupError::InvalidDescriptor {
1081 message: format!(
1082 "byte span {} is not divisible by target element size {}",
1083 span.byte_len(),
1084 target_element_size
1085 ),
1086 },
1087 ));
1088 }
1089 let layout = match TensorLayout::<DynRank>::from_parts(
1090 shape.into(),
1091 strides.into(),
1092 offset,
1093 span.byte_len() / target_element_size,
1094 ) {
1095 Ok(layout) => layout,
1096 Err(error) => {
1097 return Err((
1098 self,
1099 GroupError::InvalidDescriptor {
1100 message: error.to_string(),
1101 },
1102 ))
1103 }
1104 };
1105 let write_injective = descriptor.write_injective;
1106 let envelope = match reachable_envelope(&span, &layout, target_element_size) {
1107 Ok(envelope) => envelope,
1108 Err(error) => return Err((self, error)),
1109 };
1110 let (checked, _) = match validate_descriptor::<U, DynRank>(
1111 root,
1112 span,
1113 layout.shape().iter().copied().collect(),
1114 layout.strides().iter().copied().collect(),
1115 layout.offset(),
1116 write_injective,
1117 ) {
1118 Ok(value) => value,
1119 Err(error) => {
1120 return Err((
1121 self,
1122 GroupError::InvalidDescriptor {
1123 message: error.to_string(),
1124 },
1125 ))
1126 }
1127 };
1128 let element_count = match logical_element_count(layout.shape()) {
1129 Ok(count) => count,
1130 Err(error) => return Err((self, error)),
1131 };
1132 self.descriptors[descriptor_index] = Some(DescriptorRecord {
1133 allocation,
1134 element_count,
1135 provider: descriptor.provider,
1136 placement: descriptor.placement.clone(),
1137 envelope,
1138 write_injective,
1139 checked,
1140 });
1141 Ok(self)
1142 }
1143
1144 pub(crate) fn view<T: TensorScalar, R: TensorRank>(
1145 &self,
1146 slot: DescriptorSlot,
1147 ) -> Result<GroupReadView<'_, T, R>, GroupError> {
1148 let (descriptor_index, descriptor) = self.resolve_descriptor(slot)?;
1149 check_typed::<T, R>(descriptor)?;
1150 let owner = self
1151 .allocations
1152 .get(descriptor.allocation.index())
1153 .ok_or(GroupError::AllocationSlotOutOfBounds {
1154 slot: descriptor.allocation.index(),
1155 })?
1156 .as_ref()
1157 .ok_or(GroupError::AllocationSlotVacant {
1158 slot: descriptor.allocation.index(),
1159 })?;
1160 let _ = descriptor_index;
1161 Ok(GroupReadView {
1162 owner: NonNull::from(owner),
1163 descriptor: descriptor.clone(),
1164 _borrow: PhantomData,
1165 })
1166 }
1167
1168 pub(crate) fn view_raw<T: 'static, R: TensorRank>(
1169 &self,
1170 slot: DescriptorSlot,
1171 ) -> Result<GroupReadView<'_, T, R>, GroupError> {
1172 let (_, descriptor) = self.resolve_descriptor(slot)?;
1173 let owner = self
1174 .allocations
1175 .get(descriptor.allocation.index())
1176 .ok_or(GroupError::AllocationSlotOutOfBounds {
1177 slot: descriptor.allocation.index(),
1178 })?
1179 .as_ref()
1180 .ok_or(GroupError::AllocationSlotVacant {
1181 slot: descriptor.allocation.index(),
1182 })?;
1183 Ok(GroupReadView {
1184 owner: NonNull::from(owner),
1185 descriptor: descriptor.clone(),
1186 _borrow: PhantomData,
1187 })
1188 }
1189
1190 pub(crate) fn prepare_device_read_for_layout<T: TensorScalar, R: TensorRank>(
1191 &self,
1192 slot: DescriptorSlot,
1193 layout: &TensorLayout<R>,
1194 ) -> Result<Box<dyn crate::PreparedDeviceAccess + '_>, AccessError> {
1195 self.view_raw::<T, R>(slot)
1196 .map_err(|error| AccessError::InvalidLayout {
1197 message: error.to_string(),
1198 })?
1199 .prepare_device_read_for_layout(layout)
1200 }
1201
1202 pub(crate) fn allocation_index(
1203 &self,
1204 slot: DescriptorSlot,
1205 ) -> Result<AllocationSlot, GroupError> {
1206 Ok(self.resolve_descriptor(slot)?.1.allocation)
1207 }
1208
1209 pub(crate) fn set_host_recycler<T: TensorScalar>(
1210 &mut self,
1211 allocation_index: usize,
1212 recycler: std::sync::Weak<dyn super::root::HostBufferRecycler<T>>,
1213 ) -> Result<(), AccessError> {
1214 let owner = self
1215 .allocations
1216 .get_mut(allocation_index)
1217 .and_then(Option::as_mut)
1218 .ok_or(AccessError::Unsupported {
1219 backend: "missing host allocation",
1220 })?;
1221 owner.set_host_recycler(recycler)
1222 }
1223
1224 pub(crate) fn host_buffer_at<T: 'static>(
1225 &self,
1226 allocation_index: AllocationSlot,
1227 ) -> Option<&crate::StorageBuffer<T>> {
1228 self.allocations
1229 .get(allocation_index.index())?
1230 .as_ref()?
1231 .host_buffer::<T>()
1232 }
1233
1234 pub(crate) fn host_root_metadata<T: 'static>(
1235 &self,
1236 slot: DescriptorSlot,
1237 ) -> Option<(usize, usize)> {
1238 let (_, descriptor) = self.resolve_descriptor(slot).ok()?;
1239 let owner = self
1240 .allocations
1241 .get(descriptor.allocation.index())?
1242 .as_ref()?;
1243 if descriptor.span() != owner.root_span() {
1244 return None;
1245 }
1246 let crate::StorageBuffer::Host(data) = owner.host_buffer::<T>()? else {
1247 return None;
1248 };
1249 let pointer = data.as_ptr() as usize;
1250 let byte_len = data.len().checked_mul(size_of::<T>())?;
1251 Some((pointer, byte_len))
1252 }
1253
1254 pub(crate) fn backend_buffer<T: 'static>(
1255 &self,
1256 slot: DescriptorSlot,
1257 ) -> Option<&crate::StorageBuffer<T>> {
1258 let (_, descriptor) = self.resolve_descriptor(slot).ok()?;
1259 self.allocations
1260 .get(descriptor.allocation.index())?
1261 .as_ref()?
1262 .backend_buffer::<T>()
1263 }
1264
1265 pub(crate) fn descriptor_dtype(&self, slot: DescriptorSlot) -> Option<DType> {
1266 self.resolve_descriptor(slot)
1267 .ok()
1268 .map(|(_, descriptor)| descriptor.dtype())
1269 }
1270
1271 fn allocation_root_resource(&self, allocation: AllocationSlot) -> Option<RootResourceId> {
1277 self.allocations
1278 .get(allocation.index())?
1279 .as_ref()
1280 .map(|owner| owner.root_span().root_resource())
1281 }
1282
1283 pub(crate) fn descriptor_len(&self, slot: DescriptorSlot) -> Option<usize> {
1284 self.resolve_descriptor(slot)
1285 .ok()
1286 .map(|(_, descriptor)| descriptor.element_count)
1287 }
1288
1289 pub(crate) fn backend_identity(
1290 &self,
1291 slot: DescriptorSlot,
1292 ) -> Option<(crate::AllocationDomainId, crate::AllocationId)> {
1293 let (_, descriptor) = self.resolve_descriptor(slot).ok()?;
1294 if descriptor.provider == ProviderKind::Cpu {
1295 return None;
1296 }
1297 let owner = self
1298 .allocations
1299 .get(descriptor.allocation.index())?
1300 .as_ref()?;
1301 let key = owner.root_identity().extent().key();
1302 Some((key.domain(), key.local()))
1303 }
1304
1305 pub(crate) fn provider_kind(&self, slot: DescriptorSlot) -> Option<ProviderKind> {
1306 self.resolve_descriptor(slot)
1307 .ok()
1308 .map(|(_, descriptor)| descriptor.provider)
1309 }
1310
1311 pub(crate) fn backend_root_buffer<T: 'static>(&self) -> Option<&crate::StorageBuffer<T>> {
1312 self.allocations
1313 .first()
1314 .and_then(|owner| owner.as_ref())
1315 .and_then(|owner| owner.backend_buffer::<T>())
1316 }
1317
1318 pub(crate) fn backend_buffer_mut<T: 'static>(
1319 &mut self,
1320 slot: DescriptorSlot,
1321 ) -> Option<&mut crate::StorageBuffer<T>> {
1322 let allocation = self.resolve_descriptor(slot).ok()?.1.allocation;
1323 self.allocations
1324 .get_mut(allocation.index())?
1325 .as_mut()?
1326 .backend_buffer_mut::<T>()
1327 }
1328
1329 pub(crate) fn backend_root_buffer_mut<T: 'static>(
1330 &mut self,
1331 ) -> Option<&mut crate::StorageBuffer<T>> {
1332 self.allocations
1333 .first_mut()
1334 .and_then(|owner| owner.as_mut())
1335 .and_then(|owner| owner.backend_buffer_mut::<T>())
1336 }
1337
1338 pub(crate) fn view_mut_raw<T: 'static, R: TensorRank>(
1339 &mut self,
1340 slot: DescriptorSlot,
1341 ) -> Result<GroupWriteView<'_, T, R>, GroupError> {
1342 let (_, descriptor) = self.resolve_descriptor(slot)?;
1343 let descriptor = descriptor.clone();
1344 let owner = self
1345 .allocations
1346 .get_mut(descriptor.allocation.index())
1347 .ok_or(GroupError::AllocationSlotOutOfBounds {
1348 slot: descriptor.allocation.index(),
1349 })?
1350 .as_mut()
1351 .ok_or(GroupError::AllocationSlotVacant {
1352 slot: descriptor.allocation.index(),
1353 })?;
1354 Ok(GroupWriteView {
1355 owner: NonNull::from(owner),
1356 descriptor: descriptor.clone(),
1357 _borrow: PhantomData,
1358 })
1359 }
1360
1361 pub(crate) fn prepare_device_write_for_layout<T: TensorScalar, R: TensorRank>(
1362 &mut self,
1363 slot: DescriptorSlot,
1364 layout: &TensorLayout<R>,
1365 ) -> Result<Box<dyn crate::PreparedDeviceAccess + '_>, AccessError> {
1366 self.view_mut_raw::<T, R>(slot)
1367 .map_err(|error| AccessError::InvalidLayout {
1368 message: error.to_string(),
1369 })?
1370 .prepare_device_write_for_layout(layout)
1371 }
1372
1373 pub(crate) fn view_mut<T: TensorScalar, R: TensorRank>(
1374 &mut self,
1375 slot: DescriptorSlot,
1376 ) -> Result<GroupWriteView<'_, T, R>, GroupError> {
1377 let (descriptor_index, descriptor) = self.resolve_descriptor(slot)?;
1378 let mut descriptor = descriptor.clone();
1379 check_typed::<T, R>(&descriptor)?;
1380 if !descriptor.write_injective {
1381 descriptor
1382 .layout()
1383 .validate_mutable_no_overlap()
1384 .map_err(|error| GroupError::InvalidDescriptor {
1385 message: error.to_string(),
1386 })?;
1387 if let Some(Some(retained)) = self.descriptors.get_mut(descriptor_index) {
1388 retained.write_injective = true;
1389 }
1390 descriptor.write_injective = true;
1391 }
1392 let owner = self
1393 .allocations
1394 .get_mut(descriptor.allocation.index())
1395 .ok_or(GroupError::AllocationSlotOutOfBounds {
1396 slot: descriptor.allocation.index(),
1397 })?
1398 .as_mut()
1399 .ok_or(GroupError::AllocationSlotVacant {
1400 slot: descriptor.allocation.index(),
1401 })?;
1402 Ok(GroupWriteView {
1403 owner: NonNull::from(owner),
1404 descriptor,
1405 _borrow: PhantomData,
1406 })
1407 }
1408
1409 pub(crate) fn split_mut<T: TensorScalar, R: TensorRank>(
1410 &mut self,
1411 slots: &[DescriptorSlot],
1412 ) -> Result<Vec<GroupWriteView<'_, T, R>>, DisjointViewError> {
1413 let mut selected = Vec::with_capacity(slots.len());
1414 for &slot in slots {
1415 let (descriptor_index, descriptor) = self.resolve_descriptor(slot)?;
1416 if selected
1417 .iter()
1418 .any(|(seen, _, _): &(DescriptorSlot, usize, DescriptorRecord)| *seen == slot)
1419 {
1420 return Err(DisjointViewError::DuplicateSlot { slot: slot.index() });
1421 }
1422 check_typed::<T, R>(descriptor)?;
1423 if !descriptor.write_injective {
1424 descriptor
1425 .layout()
1426 .validate_mutable_no_overlap()
1427 .map_err(|_| DisjointViewError::NonInjective { slot: slot.index() })?;
1428 }
1429 selected.push((slot, descriptor_index, descriptor.clone()));
1430 }
1431
1432 for left in 0..selected.len() {
1433 for right in (left + 1)..selected.len() {
1434 let first = &selected[left].2;
1435 let second = &selected[right].2;
1436 if first.span().root_resource() != second.span().root_resource() {
1437 continue;
1438 }
1439 match (first.envelope, second.envelope) {
1440 (None, _) | (_, None) => {}
1441 (Some(first), Some(second)) => {
1442 if first
1443 .overlaps(second)
1444 .map_err(|_| DisjointViewError::NotProvablyDisjoint)?
1445 {
1446 return Err(DisjointViewError::PairwiseOverlap);
1447 }
1448 }
1449 }
1450 }
1451 }
1452
1453 for (_, descriptor_index, _) in &selected {
1454 if let Some(Some(descriptor)) = self.descriptors.get_mut(*descriptor_index) {
1455 descriptor.write_injective = true;
1456 }
1457 }
1458
1459 let mut children = Vec::with_capacity(selected.len());
1460 for (_, _, mut descriptor) in selected {
1461 descriptor.write_injective = true;
1462 let owner = self
1463 .allocations
1464 .get_mut(descriptor.allocation.index())
1465 .ok_or(GroupError::AllocationSlotOutOfBounds {
1466 slot: descriptor.allocation.index(),
1467 })?
1468 .as_mut()
1469 .ok_or(GroupError::AllocationSlotVacant {
1470 slot: descriptor.allocation.index(),
1471 })?;
1472 children.push(GroupWriteView {
1473 owner: NonNull::from(owner),
1474 descriptor,
1475 _borrow: PhantomData,
1476 });
1477 }
1478 Ok(children)
1479 }
1480
1481 pub(crate) fn try_extract(
1482 &mut self,
1483 slot: DescriptorSlot,
1484 ) -> Result<OwnedStorage, ExtractError> {
1485 let (descriptor_index, descriptor) = self.resolve_descriptor(slot)?;
1486 let allocation = descriptor.allocation;
1487 let references = self
1488 .descriptors
1489 .iter()
1490 .flatten()
1491 .filter(|candidate| candidate.allocation == allocation)
1492 .count();
1493 if references != 1 {
1494 return Err(ExtractError::AliasedAllocation {
1495 allocation: allocation.index(),
1496 });
1497 }
1498 let _ = self.descriptors[descriptor_index].take();
1499 self.allocations
1500 .get_mut(allocation.index())
1501 .ok_or(GroupError::AllocationSlotOutOfBounds {
1502 slot: allocation.index(),
1503 })?
1504 .take()
1505 .ok_or(GroupError::AllocationSlotVacant {
1506 slot: allocation.index(),
1507 })
1508 .map_err(ExtractError::from)
1509 }
1510
1511 #[allow(clippy::result_large_err)]
1524 pub fn take_tensor(&mut self, slot: DescriptorSlot) -> Result<crate::Tensor, GroupError> {
1525 let (_, descriptor) = self.resolve_descriptor(slot)?;
1526 self.ensure_whole_compact(slot, descriptor)?;
1527 let descriptor = descriptor.clone();
1528 let dtype = descriptor.dtype();
1529 let layout = descriptor.layout().clone();
1530 let placement = descriptor.placement.clone();
1531 let owner = self
1532 .try_extract(slot)
1533 .map_err(|error| GroupError::InvalidDescriptor {
1534 message: error.to_string(),
1535 })?;
1536 let mut extracted = Self::new();
1537 extracted.allocations.push(Some(owner));
1538 let mut descriptor = descriptor.clone();
1539 descriptor.allocation = AllocationSlot(0);
1540 extracted.descriptors.push(Some(descriptor));
1541 Ok(tensor_from_group(
1542 extracted,
1543 DescriptorSlot(0),
1544 AllocationSlot(0),
1545 dtype,
1546 layout,
1547 placement,
1548 ))
1549 }
1550
1551 #[allow(clippy::result_large_err)]
1554 pub(crate) fn into_owner(
1555 mut self,
1556 slot: DescriptorSlot,
1557 ) -> Result<OwnedStorage, (Self, ExtractError)> {
1558 let result = self.try_extract(slot);
1559 match result {
1560 Ok(owner) => Ok(owner),
1561 Err(error) => Err((self, error)),
1562 }
1563 }
1564
1565 #[allow(clippy::result_large_err)]
1581 pub fn into_tensor(self, slot: DescriptorSlot) -> Result<crate::Tensor, (Self, GroupError)> {
1582 let (_, descriptor) = match self.resolve_descriptor(slot) {
1583 Ok(value) => value,
1584 Err(error) => return Err((self, error)),
1585 };
1586 let allocation = descriptor.allocation;
1587 let references = self
1588 .descriptors
1589 .iter()
1590 .flatten()
1591 .filter(|candidate| candidate.allocation == allocation)
1592 .count();
1593 if references != 1 {
1594 return Err((
1595 self,
1596 GroupError::AliasedAllocation {
1597 allocation: allocation.index(),
1598 },
1599 ));
1600 }
1601 if let Err(error) = self.ensure_whole_compact(slot, descriptor) {
1602 return Err((self, error));
1603 }
1604 let dtype = descriptor.dtype();
1605 let layout = descriptor.layout().clone();
1606 let placement = descriptor.placement.clone();
1607 let (group, slot) = match self.into_single_descriptor(slot) {
1608 Ok(value) => value,
1609 Err((group, error)) => return Err((group, error)),
1610 };
1611 Ok(tensor_from_group(
1612 group,
1613 slot,
1614 AllocationSlot(0),
1615 dtype,
1616 layout,
1617 placement,
1618 ))
1619 }
1620
1621 #[allow(clippy::result_large_err)]
1624 fn into_single_descriptor(
1625 mut self,
1626 slot: DescriptorSlot,
1627 ) -> Result<(Self, DescriptorSlot), (Self, GroupError)> {
1628 let (descriptor_index, descriptor) = match self.resolve_descriptor(slot) {
1629 Ok(value) => value,
1630 Err(error) => return Err((self, error)),
1631 };
1632 let allocation = descriptor.allocation;
1633 let mut descriptor = match self.descriptors[descriptor_index].take() {
1634 Some(descriptor) => descriptor,
1635 None => {
1636 return Err((
1637 self,
1638 GroupError::DescriptorSlotVacant { slot: slot.index() },
1639 ))
1640 }
1641 };
1642 let owner = match self.allocations.get_mut(allocation.index()) {
1643 Some(owner) => match owner.take() {
1644 Some(owner) => owner,
1645 None => {
1646 self.descriptors[descriptor_index] = Some(descriptor);
1647 return Err((
1648 self,
1649 GroupError::AllocationSlotVacant {
1650 slot: allocation.index(),
1651 },
1652 ));
1653 }
1654 },
1655 None => {
1656 self.descriptors[descriptor_index] = Some(descriptor);
1657 return Err((
1658 self,
1659 GroupError::AllocationSlotOutOfBounds {
1660 slot: allocation.index(),
1661 },
1662 ));
1663 }
1664 };
1665 let mut group = Self::new();
1666 descriptor.allocation = AllocationSlot(0);
1667 group.allocations.push(Some(owner));
1668 group.descriptors.push(Some(descriptor));
1669 Ok((group, DescriptorSlot(0)))
1670 }
1671
1672 #[allow(clippy::result_large_err)]
1678 pub(crate) fn into_host_vec<T: 'static>(
1679 mut self,
1680 slot: DescriptorSlot,
1681 ) -> Result<Vec<T>, (Self, String)> {
1682 let allocation = match self.resolve_descriptor(slot) {
1686 Ok((_, descriptor)) => descriptor.allocation,
1687 Err(error) => return Err((self, error.to_string())),
1688 };
1689 let owner = match self
1690 .allocations
1691 .get_mut(allocation.index())
1692 .map(Option::take)
1693 {
1694 Some(Some(owner)) => owner,
1695 Some(None) => {
1696 return Err((
1697 self,
1698 GroupError::AllocationSlotVacant {
1699 slot: allocation.index(),
1700 }
1701 .to_string(),
1702 ))
1703 }
1704 None => {
1705 return Err((
1706 self,
1707 GroupError::AllocationSlotOutOfBounds {
1708 slot: allocation.index(),
1709 }
1710 .to_string(),
1711 ))
1712 }
1713 };
1714 match owner.into_host_vec::<T>() {
1715 Ok(data) => Ok(data),
1716 Err((owner, error)) => {
1717 self.allocations[allocation.index()] = Some(owner);
1718 Err((self, error.to_string()))
1719 }
1720 }
1721 }
1722
1723 fn ensure_whole_compact(
1728 &self,
1729 slot: DescriptorSlot,
1730 descriptor: &DescriptorRecord,
1731 ) -> Result<(), GroupError> {
1732 let index = descriptor.allocation.index();
1733 let owner = self
1734 .allocations
1735 .get(index)
1736 .ok_or(GroupError::AllocationSlotOutOfBounds { slot: index })?
1737 .as_ref()
1738 .ok_or(GroupError::AllocationSlotVacant { slot: index })?;
1739 let root = owner.as_ref().root_identity().root_span();
1740 let span = descriptor.span();
1741 let layout = descriptor.layout();
1742 let compact =
1743 layout
1744 .is_compact_col_major()
1745 .map_err(|error| GroupError::InvalidDescriptor {
1746 message: error.to_string(),
1747 })?;
1748 let whole = span.byte_offset() == root.byte_offset()
1749 && span.byte_len() == root.byte_len()
1750 && descriptor
1751 .element_count()
1752 .checked_mul(descriptor.element_size())
1753 == Some(span.byte_len());
1754 if compact && layout.offset() == 0 && whole {
1755 Ok(())
1756 } else {
1757 Err(GroupError::NonCompactDescriptor { slot: slot.index() })
1758 }
1759 }
1760
1761 fn resolve_descriptor(
1762 &self,
1763 slot: DescriptorSlot,
1764 ) -> Result<(usize, &DescriptorRecord), GroupError> {
1765 let index = slot.index();
1766 let descriptor = self
1767 .descriptors
1768 .get(index)
1769 .ok_or(GroupError::DescriptorSlotOutOfBounds { slot: index })?
1770 .as_ref()
1771 .ok_or(GroupError::DescriptorSlotVacant { slot: index })?;
1772 Ok((index, descriptor))
1773 }
1774
1775 #[cfg(test)]
1776 pub(crate) fn test_vacate_allocation(&mut self, slot: AllocationSlot) {
1777 if let Some(entry) = self.allocations.get_mut(slot.index()) {
1778 *entry = None;
1779 }
1780 }
1781}
1782
1783fn check_typed<T: TensorScalar, R: TensorRank>(
1784 descriptor: &DescriptorRecord,
1785) -> Result<(), GroupError> {
1786 if descriptor.dtype() != T::dtype() {
1787 return Err(GroupError::DTypeMismatch {
1788 expected: descriptor.dtype(),
1789 actual: T::dtype(),
1790 });
1791 }
1792 if let Some(expected) = R::RANK {
1793 let actual = descriptor.layout().shape().len();
1794 if expected != actual {
1795 return Err(GroupError::RankMismatch { expected, actual });
1796 }
1797 }
1798 Ok(())
1799}
1800
1801fn logical_element_count(shape: &[usize]) -> Result<usize, GroupError> {
1802 shape.iter().try_fold(1usize, |count, &extent| {
1803 count
1804 .checked_mul(extent)
1805 .ok_or_else(|| GroupError::InvalidDescriptor {
1806 message: "logical element count overflows".to_owned(),
1807 })
1808 })
1809}
1810
1811fn reachable_envelope(
1812 span: &RootBoundSpan,
1813 layout: &TensorLayout<DynRank>,
1814 element_size: usize,
1815) -> Result<Option<ByteRange>, GroupError> {
1816 if layout.shape().contains(&0) {
1817 return Ok(None);
1818 }
1819 let mut minimum = layout.offset() as i128;
1820 let mut maximum = minimum;
1821 for (&extent, &stride) in layout.shape().iter().zip(layout.strides()) {
1822 let steps = i128::try_from(extent - 1).map_err(|_| GroupError::InvalidDescriptor {
1823 message: "layout extent does not fit i128".to_owned(),
1824 })?;
1825 let contribution =
1826 (stride as i128)
1827 .checked_mul(steps)
1828 .ok_or_else(|| GroupError::InvalidDescriptor {
1829 message: "reachable layout arithmetic overflows".to_owned(),
1830 })?;
1831 if contribution < 0 {
1832 minimum =
1833 minimum
1834 .checked_add(contribution)
1835 .ok_or_else(|| GroupError::InvalidDescriptor {
1836 message: "reachable layout minimum overflows".to_owned(),
1837 })?;
1838 } else {
1839 maximum =
1840 maximum
1841 .checked_add(contribution)
1842 .ok_or_else(|| GroupError::InvalidDescriptor {
1843 message: "reachable layout maximum overflows".to_owned(),
1844 })?;
1845 }
1846 }
1847 let minimum = usize::try_from(minimum).map_err(|_| GroupError::InvalidDescriptor {
1848 message: "reachable layout minimum is negative".to_owned(),
1849 })?;
1850 let maximum = usize::try_from(maximum).map_err(|_| GroupError::InvalidDescriptor {
1851 message: "reachable layout maximum is negative or too large".to_owned(),
1852 })?;
1853 let byte_offset = span
1854 .byte_offset()
1855 .checked_add(minimum.checked_mul(element_size).ok_or_else(|| {
1856 GroupError::InvalidDescriptor {
1857 message: "reachable byte offset overflows".to_owned(),
1858 }
1859 })?)
1860 .ok_or_else(|| GroupError::InvalidDescriptor {
1861 message: "reachable byte offset overflows".to_owned(),
1862 })?;
1863 let byte_len = maximum
1864 .checked_sub(minimum)
1865 .and_then(|length| length.checked_add(1))
1866 .and_then(|length| length.checked_mul(element_size))
1867 .ok_or_else(|| GroupError::InvalidDescriptor {
1868 message: "reachable byte length overflows".to_owned(),
1869 })?;
1870 let range = ByteRange::new(byte_offset, byte_len);
1871 range
1872 .checked_end()
1873 .map_err(|error| GroupError::InvalidDescriptor {
1874 message: error.to_string(),
1875 })?;
1876 Ok(Some(range))
1877}
1878
1879#[cfg(test)]
1880pub(crate) fn test_logical_element_count(shape: &[usize]) -> Result<usize, GroupError> {
1881 logical_element_count(shape)
1882}
1883
1884#[cfg(test)]
1885pub(crate) fn test_reachable_envelope(
1886 span: RootBoundSpan,
1887 layout: TensorLayout<DynRank>,
1888 element_size: usize,
1889) -> Result<Option<ByteRange>, GroupError> {
1890 reachable_envelope(&span, &layout, element_size)
1891}