1use tenferro_tensor::{
2 DType, DotGeneralAccumulation, DotGeneralConfig, ShapeMismatch, Tensor, TensorRead, TensorView,
3 TensorViewMut, TensorWrite, TypedTensor, ValidationError,
4};
5
6use num_complex::{Complex32, Complex64};
7use smallvec::SmallVec;
8use std::mem::MaybeUninit;
9use std::sync::atomic::{AtomicUsize, Ordering};
10use std::sync::Arc;
11
12use crate::backend::CpuBackendKind;
13use crate::batch_policy::strategy_unavailable;
14use crate::buffer_pool::{BufferPool, PoolScalar};
15use crate::provider::{
16 builtin_gemm_provider, builtin_layout_provider, CpuContractionAxes, CpuDotGeneralRequest,
17 CpuExecutionContext, CpuGemmProvider, CpuGeneralContractionProvider, CpuGroupedGemmRequest,
18 CpuLayoutTransformIntent, CpuLayoutTransformProvider, CpuLayoutTransformRequest,
19 CpuOperationEntry, CpuProviderOutcome, CpuProviderUnsupported, CpuUninitGemmProvider,
20};
21use crate::CpuBatchStrategy;
22use crate::{
23 gemm::GemmAnalysisCache, CpuDomainExecutorError, CpuDomainId, CpuProviderDomainError, Error,
24 ParallelMode, PooledUninitOutput, Result,
25};
26
27const OP: &str = "dot_general";
28
29#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
42pub enum GeneralContractionPolicy {
43 #[default]
45 Preferred,
46 Required,
48}
49
50#[derive(Debug)]
51pub(crate) struct DotGeneralRuntime {
52 pub(crate) general: Option<Arc<dyn CpuGeneralContractionProvider>>,
53 pub(crate) gemm: Arc<dyn CpuGemmProvider>,
54 pub(crate) layout: Arc<dyn CpuLayoutTransformProvider>,
55 general_capabilities: Option<crate::CpuProviderExecutionCapabilities>,
56 gemm_capabilities: crate::CpuProviderExecutionCapabilities,
57 layout_capabilities: crate::CpuProviderExecutionCapabilities,
58 pub(crate) general_policy: GeneralContractionPolicy,
59 grouped_scheduling: GroupedGemmScheduling,
60 capability_policy: ProviderCapabilityPolicy,
61}
62
63#[derive(Clone, Copy, Debug, PartialEq, Eq)]
64enum GroupedGemmScheduling {
65 ProviderOwned,
66 EngineOuter,
67}
68
69#[derive(Clone, Copy, Debug, PartialEq, Eq)]
70enum ProviderCapabilityPolicy {
71 Strict,
72 ProviderDefaultCompatibility,
73}
74
75const GROUPED_JOB_STATE_BITS: usize = 2;
76const GROUPED_JOBS_PER_STATE_WORD: usize = usize::BITS as usize / GROUPED_JOB_STATE_BITS;
77const GROUPED_INLINE_STATE_WORDS: usize = 4;
78#[cfg(test)]
79const GROUPED_INLINE_JOB_CAPACITY: usize = GROUPED_INLINE_STATE_WORDS * GROUPED_JOBS_PER_STATE_WORD;
80
81#[derive(Clone, Copy, Debug, Eq, PartialEq)]
82#[repr(usize)]
83enum GroupedJobState {
84 Unclaimed = 0,
85 Running = 1,
86 Complete = 2,
87 Reserved = 3,
88}
89
90impl GroupedJobState {
91 fn from_bits(bits: usize) -> Self {
92 match bits {
93 0 => Self::Unclaimed,
94 1 => Self::Running,
95 2 => Self::Complete,
96 _ => Self::Reserved,
97 }
98 }
99}
100
101struct PackedJobStates {
107 words: SmallVec<[AtomicUsize; GROUPED_INLINE_STATE_WORDS]>,
108 len: usize,
109}
110
111impl PackedJobStates {
112 fn new(len: usize) -> Self {
113 let word_count = len.div_ceil(GROUPED_JOBS_PER_STATE_WORD);
114 let mut words = SmallVec::new();
115 words.resize_with(word_count, || AtomicUsize::new(0));
116 Self { words, len }
117 }
118
119 fn position(index: usize) -> (usize, usize) {
120 let word = index / GROUPED_JOBS_PER_STATE_WORD;
121 let shift = (index % GROUPED_JOBS_PER_STATE_WORD) * GROUPED_JOB_STATE_BITS;
122 (word, shift)
123 }
124
125 fn state(&self, index: usize) -> GroupedJobState {
126 let (word, shift) = Self::position(index);
127 let bits = (self.words[word].load(Ordering::Acquire) >> shift) & 0b11;
128 GroupedJobState::from_bits(bits)
129 }
130
131 fn try_claim(&self, index: usize) -> std::result::Result<(), GroupedJobState> {
132 let (word, shift) = Self::position(index);
133 let word = &self.words[word];
134 let mask = 0b11usize << shift;
135 let running = (GroupedJobState::Running as usize) << shift;
136 let mut observed = word.load(Ordering::Acquire);
137 loop {
138 let state = GroupedJobState::from_bits((observed & mask) >> shift);
139 if state != GroupedJobState::Unclaimed {
140 return Err(state);
141 }
142 let updated = (observed & !mask) | running;
143 match word.compare_exchange_weak(observed, updated, Ordering::AcqRel, Ordering::Acquire)
144 {
145 Ok(_) => return Ok(()),
146 Err(current) => observed = current,
147 }
148 }
149 }
150
151 fn complete(&self, index: usize) -> bool {
152 let (word, shift) = Self::position(index);
153 let word = &self.words[word];
154 let mask = 0b11usize << shift;
155 let complete = (GroupedJobState::Complete as usize) << shift;
156 let mut observed = word.load(Ordering::Acquire);
157 loop {
158 if GroupedJobState::from_bits((observed & mask) >> shift) != GroupedJobState::Running {
159 return false;
160 }
161 let updated = (observed & !mask) | complete;
162 match word.compare_exchange_weak(observed, updated, Ordering::AcqRel, Ordering::Acquire)
163 {
164 Ok(_) => return true,
165 Err(current) => observed = current,
166 }
167 }
168 }
169
170 fn first_incomplete(&self) -> Option<(usize, GroupedJobState)> {
171 (0..self.len)
172 .map(|index| (index, self.state(index)))
173 .find(|(_, state)| *state != GroupedJobState::Complete)
174 }
175
176 #[cfg(test)]
177 fn len(&self) -> usize {
178 self.len
179 }
180
181 #[cfg(test)]
182 fn word_count(&self) -> usize {
183 self.words.len()
184 }
185
186 #[cfg(test)]
187 fn spilled(&self) -> bool {
188 self.words.spilled()
189 }
190}
191
192fn standard_grouped_scheduling(kind: CpuBackendKind) -> GroupedGemmScheduling {
193 match kind {
194 CpuBackendKind::Faer => GroupedGemmScheduling::EngineOuter,
195 CpuBackendKind::Blas => GroupedGemmScheduling::ProviderOwned,
196 }
197}
198
199#[derive(Debug)]
200pub(crate) struct CpuProviderBundleInner {
201 pub(crate) dot_general: DotGeneralRuntime,
202 pub(crate) extensions: crate::provider_extensions::ProviderExtensions,
204}
205
206#[derive(Clone, Copy, Debug)]
207pub(crate) enum CpuProviderDomainContract {
208 CooperativeCpuSet,
209 CallerManaged,
210}
211
212#[derive(Clone, Debug)]
227pub struct CpuProviderBundle {
228 inner: Arc<CpuProviderBundleInner>,
229}
230
231impl CpuProviderBundle {
232 pub(crate) fn standard(kind: CpuBackendKind, provider_default_compatibility: bool) -> Self {
233 let gemm = builtin_gemm_provider(kind);
234 let layout = builtin_layout_provider();
235 let gemm_capabilities = gemm.execution_capabilities();
236 let layout_capabilities = layout.execution_capabilities();
237 Self {
238 inner: Arc::new(CpuProviderBundleInner {
239 dot_general: DotGeneralRuntime {
240 general: None,
241 gemm,
242 layout,
243 general_capabilities: None,
244 gemm_capabilities,
245 layout_capabilities,
246 general_policy: GeneralContractionPolicy::Preferred,
247 grouped_scheduling: standard_grouped_scheduling(kind),
248 capability_policy: if provider_default_compatibility {
249 ProviderCapabilityPolicy::ProviderDefaultCompatibility
250 } else {
251 ProviderCapabilityPolicy::Strict
252 },
253 },
254 extensions: crate::provider_extensions::ProviderExtensions::default(),
255 }),
256 }
257 }
258
259 pub fn builder(kind: CpuBackendKind) -> CpuProviderBundleBuilder {
261 CpuProviderBundleBuilder {
262 gemm: Some(builtin_gemm_provider(kind)),
263 layout: Some(builtin_layout_provider()),
264 general: None,
265 general_policy: GeneralContractionPolicy::Preferred,
266 grouped_scheduling: standard_grouped_scheduling(kind),
267 capability_policy: ProviderCapabilityPolicy::Strict,
268 extensions: crate::provider_extensions::ProviderExtensions::default(),
269 }
270 }
271
272 pub fn custom_builder() -> CpuProviderBundleBuilder {
274 CpuProviderBundleBuilder {
275 gemm: None,
276 layout: None,
277 general: None,
278 general_policy: GeneralContractionPolicy::Preferred,
279 grouped_scheduling: GroupedGemmScheduling::ProviderOwned,
280 capability_policy: ProviderCapabilityPolicy::Strict,
281 extensions: crate::provider_extensions::ProviderExtensions::default(),
282 }
283 }
284
285 pub fn extension<E: std::any::Any + Send + Sync>(&self) -> Option<Arc<E>> {
302 self.inner.extensions.get::<E>()
303 }
304
305 pub fn shares_identity_with(&self, other: &Self) -> bool {
307 Arc::ptr_eq(&self.inner, &other.inner)
308 }
309
310 pub(crate) fn inner(&self) -> &Arc<CpuProviderBundleInner> {
311 &self.inner
312 }
313
314 pub(crate) fn dot_general(&self) -> &DotGeneralRuntime {
315 &self.inner.dot_general
316 }
317
318 pub(crate) fn validate_for_domain(
319 &self,
320 domain_id: CpuDomainId,
321 thread_budget: std::num::NonZeroUsize,
322 contract: CpuProviderDomainContract,
323 ) -> std::result::Result<(), CpuProviderBundleInstallError> {
324 let runtime = self.dot_general();
325 let validate = |provider, capabilities| {
326 let result = match contract {
327 CpuProviderDomainContract::CooperativeCpuSet => {
328 crate::provider_capability::validate_provider_for_domain(
329 capabilities,
330 thread_budget,
331 )
332 }
333 CpuProviderDomainContract::CallerManaged => {
334 crate::provider_capability::validate_provider_for_caller_managed_domain(
335 capabilities,
336 thread_budget,
337 )
338 }
339 };
340 result.map_err(|source| CpuProviderBundleInstallError::IncompatibleDomain {
341 domain_id,
342 provider,
343 source,
344 })
345 };
346
347 if let Some(capabilities) = runtime.general_capabilities {
348 validate(CpuProviderSlot::GeneralContraction, capabilities)?;
349 }
350 validate(CpuProviderSlot::Gemm, runtime.gemm_capabilities)?;
351 validate(
352 CpuProviderSlot::LayoutTransform,
353 runtime.layout_capabilities,
354 )?;
355
356 let selected_mode = if thread_budget.get() == 1 {
357 ParallelMode::Sequential
358 } else if runtime.accepts_dot_general_mode(ParallelMode::Inner) {
359 ParallelMode::Inner
360 } else {
361 ParallelMode::Sequential
362 };
363 for (provider, capabilities) in [
364 (CpuProviderSlot::Gemm, runtime.gemm_capabilities),
365 (
366 CpuProviderSlot::LayoutTransform,
367 runtime.layout_capabilities,
368 ),
369 ] {
370 if !capabilities.accepts_mode(selected_mode) {
371 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
372 domain_id,
373 provider,
374 source: CpuProviderDomainError::ParallelModeNotSupported {
375 mode: selected_mode,
376 },
377 });
378 }
379 }
380 if let Some(capabilities) = runtime.general_capabilities {
381 if !capabilities.accepts_mode(selected_mode) {
382 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
383 domain_id,
384 provider: CpuProviderSlot::GeneralContraction,
385 source: CpuProviderDomainError::ParallelModeNotSupported {
386 mode: selected_mode,
387 },
388 });
389 }
390 }
391 if runtime.grouped_scheduling == GroupedGemmScheduling::EngineOuter
392 && !runtime.gemm_capabilities.accepts_mode(ParallelMode::Outer)
393 {
394 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
395 domain_id,
396 provider: CpuProviderSlot::Gemm,
397 source: CpuProviderDomainError::ParallelModeNotSupported {
398 mode: ParallelMode::Outer,
399 },
400 });
401 }
402 Ok(())
403 }
404
405 pub(crate) fn preflight_dot_general(&self, entry: &CpuOperationEntry<'_>) -> Result<()> {
406 self.inner
407 .dot_general
408 .dot_general_mode(entry, None)
409 .map(|_| ())
410 .map_err(|error| Error::backend_source(OP, error))
411 }
412
413 #[cfg(test)]
414 #[allow(clippy::too_many_arguments)]
415 pub(crate) fn execute_dot_general_into(
416 &self,
417 entry: &CpuOperationEntry<'_>,
418 buffers: &mut BufferPool,
419 cache: &mut GemmAnalysisCache,
420 cache_slot: Option<usize>,
421 lhs: TensorRead<'_>,
422 rhs: TensorRead<'_>,
423 config: &DotGeneralConfig,
424 accumulation: DotGeneralAccumulation,
425 output: TensorWrite<'_>,
426 ) -> Result<()> {
427 self.execute_dot_general_into_scoped(
428 entry,
429 None,
430 buffers,
431 cache,
432 cache_slot,
433 lhs,
434 rhs,
435 config,
436 accumulation,
437 output,
438 )
439 }
440
441 #[allow(clippy::too_many_arguments)]
442 pub(crate) fn execute_dot_general_into_scoped(
443 &self,
444 entry: &CpuOperationEntry<'_>,
445 entered: Option<&CpuExecutionContext<'_>>,
446 buffers: &mut BufferPool,
447 cache: &mut GemmAnalysisCache,
448 cache_slot: Option<usize>,
449 lhs: TensorRead<'_>,
450 rhs: TensorRead<'_>,
451 config: &DotGeneralConfig,
452 accumulation: DotGeneralAccumulation,
453 output: TensorWrite<'_>,
454 ) -> Result<()> {
455 self.inner.dot_general.execute_into(
456 &self.inner,
457 entry,
458 entered,
459 buffers,
460 cache,
461 cache_slot,
462 lhs,
463 rhs,
464 config,
465 accumulation,
466 output,
467 )
468 }
469
470 #[cfg(test)]
471 pub(crate) fn execute_grouped_gemm(
472 &self,
473 entry: &CpuOperationEntry<'_>,
474 lhs: TensorRead<'_>,
475 rhs: TensorRead<'_>,
476 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
477 output: TensorWrite<'_>,
478 ) -> Result<()> {
479 self.execute_grouped_gemm_scoped(entry, None, lhs, rhs, config, output)
480 }
481
482 pub(crate) fn execute_grouped_gemm_scoped(
483 &self,
484 entry: &CpuOperationEntry<'_>,
485 entered: Option<&CpuExecutionContext<'_>>,
486 lhs: TensorRead<'_>,
487 rhs: TensorRead<'_>,
488 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
489 output: TensorWrite<'_>,
490 ) -> Result<()> {
491 self.inner
492 .dot_general
493 .execute_grouped(entry, entered, lhs, rhs, config, output)
494 }
495}
496
497fn execute_all_batch_elementwise(
500 context: &CpuExecutionContext<'_>,
501 buffers: &mut BufferPool,
502 lhs: &TensorRead<'_>,
503 rhs: &TensorRead<'_>,
504 config: &DotGeneralConfig,
505 accumulation: DotGeneralAccumulation,
506 output: TensorWrite<'_>,
507) -> Result<()> {
508 let exec = context.strided_exec_context();
509 let lhs_view = TensorRead::from_view(transposed_read_view(lhs, &config.lhs_batch_dims)?);
512 let rhs_view = TensorRead::from_view(transposed_read_view(rhs, &config.rhs_batch_dims)?);
513 let lhs_conj = if accumulation.lhs_conj {
514 Some(crate::elementwise::conj_read_with_pool(
515 buffers,
516 &exec,
517 lhs_view.clone(),
518 )?)
519 } else {
520 None
521 };
522 let rhs_conj = if accumulation.rhs_conj {
523 Some(crate::elementwise::conj_read_with_pool(
524 buffers,
525 &exec,
526 rhs_view.clone(),
527 )?)
528 } else {
529 None
530 };
531 let factors = [
532 lhs_conj.as_ref().map_or(lhs_view, TensorRead::from_tensor),
533 rhs_conj.as_ref().map_or(rhs_view, TensorRead::from_tensor),
534 ];
535 let overwrite = DotGeneralAccumulation::overwrite(lhs.dtype())?;
536 let (result, product) =
537 if accumulation.alpha == overwrite.alpha && accumulation.beta == overwrite.beta {
538 let result = tenferro_internal_cpu_kernels::elementwise_read_into_with_context(
540 tenferro_tensor::ElementwiseReadOp::Multiply,
541 &factors,
542 output,
543 &exec,
544 |inputs, out| {
545 crate::backend::elementwise_read_into_fallback_with_pool(
546 buffers,
547 &exec,
548 tenferro_tensor::ElementwiseReadOp::Multiply,
549 inputs,
550 out,
551 )
552 },
553 );
554 (result, None)
555 } else {
556 let [lhs_factor, rhs_factor] = factors;
557 let product =
558 crate::elementwise::mul_read_with_pool(buffers, &exec, lhs_factor, rhs_factor)?;
559 let result = crate::blas1::axpby_read_into_accum(
560 context,
561 buffers,
562 accumulation.alpha,
563 TensorRead::from_tensor(&product),
564 accumulation.beta,
565 output,
566 );
567 (result, Some(product))
568 };
569 for temporary in [lhs_conj, rhs_conj, product].into_iter().flatten() {
570 crate::backend::reclaim_tensor(buffers, temporary);
571 }
572 result
573}
574
575#[allow(clippy::large_enum_variant)]
579enum CanonicalOperand<'input> {
580 Borrowed(TensorRead<'input>),
581 Packed(Tensor),
582}
583
584impl CanonicalOperand<'_> {
585 fn read(&self) -> TensorRead<'_> {
586 match self {
587 Self::Borrowed(read) => read.clone(),
588 Self::Packed(tensor) => TensorRead::from_tensor(tensor),
589 }
590 }
591
592 fn reclaim(self, buffers: &mut BufferPool) {
593 if let Self::Packed(tensor) = self {
594 crate::backend::reclaim_tensor(buffers, tensor);
595 }
596 }
597}
598
599fn is_all_batch_contraction(
601 lhs: &TensorRead<'_>,
602 rhs: &TensorRead<'_>,
603 config: &DotGeneralConfig,
604) -> bool {
605 lhs.shape().len() == config.lhs_batch_dims.len()
606 && rhs.shape().len() == config.rhs_batch_dims.len()
607}
608
609fn dot_batch_items(lhs: &TensorRead<'_>, config: &DotGeneralConfig) -> Result<usize> {
611 config
612 .lhs_batch_dims
613 .iter()
614 .try_fold(1usize, |items, &axis| {
615 lhs.shape()
616 .get(axis)
617 .and_then(|&extent| items.checked_mul(extent))
618 })
619 .ok_or_else(|| {
620 Error::invalid_argument(OP, "lhs", "batch axes are out of range or overflow usize")
621 })
622}
623
624fn effective_batch_policy(
627 entry: &CpuOperationEntry<'_>,
628 entered: Option<&CpuExecutionContext<'_>>,
629) -> crate::CpuBatchPolicy {
630 entered.map_or_else(|| entry.batch_policy(), CpuExecutionContext::batch_policy)
631}
632
633fn vendor_batch_for(
635 policy: crate::CpuBatchPolicy,
636 strategy: CpuBatchStrategy,
637) -> crate::provider::CpuVendorBatch {
638 match strategy {
639 CpuBatchStrategy::WholeBatchVendor => crate::provider::CpuVendorBatch::Required,
640 CpuBatchStrategy::Auto => crate::provider::CpuVendorBatch::Allowed {
641 max_item_dim: policy.thresholds().vendor_batch_max_item_dim(),
642 },
643 _ => crate::provider::CpuVendorBatch::Forbidden,
644 }
645}
646
647fn unsupported_provider_error(capability: &'static str, reason: CpuProviderUnsupported) -> Error {
648 Error::unsupported(
649 OP,
650 format!("configured CPU {capability} provider reported unsupported: {reason:?}"),
651 )
652}
653
654impl DotGeneralRuntime {
655 fn accepts_dot_general_mode(&self, mode: crate::ParallelMode) -> bool {
656 self.general_capabilities
657 .is_none_or(|capabilities| capabilities.accepts_mode(mode))
658 && self.gemm_capabilities.accepts_mode(mode)
659 && self.layout_capabilities.accepts_mode(mode)
660 }
661
662 fn validate_strict_capability(
663 &self,
664 capabilities: crate::CpuProviderExecutionCapabilities,
665 thread_budget: usize,
666 ) -> std::result::Result<(), CpuProviderDomainError> {
667 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
668 return Ok(());
669 }
670 if capabilities.thread_count == crate::CpuThreadCountControl::GlobalOrUncontrolled {
671 return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
672 thread_budget,
673 control: capabilities.thread_count,
674 });
675 }
676 Ok(())
677 }
678
679 fn dot_general_mode(
680 &self,
681 entry: &CpuOperationEntry<'_>,
682 entered: Option<&CpuExecutionContext<'_>>,
683 ) -> std::result::Result<ParallelMode, CpuProviderDomainError> {
684 if entered.is_some_and(CpuExecutionContext::is_outer_fan_out_lane) {
688 crate::provider::check_outer_fan_out_delegates(
689 self.general_capabilities
690 .iter()
691 .chain([&self.gemm_capabilities, &self.layout_capabilities]),
692 )?;
693 return Ok(ParallelMode::Sequential);
694 }
695 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
696 return Ok(entry.provider_default_compatibility_mode());
697 }
698 let thread_budget = entry.thread_budget().get();
699 if let Some(capabilities) = self.general_capabilities {
700 self.validate_strict_capability(capabilities, thread_budget)?;
701 }
702 self.validate_strict_capability(self.gemm_capabilities, thread_budget)?;
703 self.validate_strict_capability(self.layout_capabilities, thread_budget)?;
704 entry.preferred_provider_mode(|mode| self.accepts_dot_general_mode(mode))
705 }
706
707 fn batched_dot_mode(
710 &self,
711 entry: &CpuOperationEntry<'_>,
712 entered: Option<&CpuExecutionContext<'_>>,
713 mode: ParallelMode,
714 items: usize,
715 ) -> Result<ParallelMode> {
716 if items <= 1 {
717 return Ok(mode);
718 }
719 let strategy = effective_batch_policy(entry, entered).strategy();
720 match strategy {
721 CpuBatchStrategy::Sequential => {
722 if self.accepts_dot_general_mode(ParallelMode::Sequential) {
723 Ok(ParallelMode::Sequential)
724 } else {
725 Err(strategy_unavailable(
726 OP,
727 strategy,
728 "a contraction provider does not accept sequential calls",
729 ))
730 }
731 }
732 CpuBatchStrategy::OuterParallel if entry.thread_budget().get() > 1 => {
736 Ok(ParallelMode::Inner)
737 }
738 CpuBatchStrategy::OuterParallel => Err(strategy_unavailable(
739 OP,
740 strategy,
741 "the selected CPU domain has one thread, so there are no outer lanes",
742 )),
743 _ => Ok(mode),
744 }
745 }
746
747 fn grouped_mode(
748 &self,
749 entry: &CpuOperationEntry<'_>,
750 entered: Option<&CpuExecutionContext<'_>>,
751 ) -> std::result::Result<ParallelMode, CpuProviderDomainError> {
752 if entered.is_some_and(CpuExecutionContext::is_outer_fan_out_lane) {
753 crate::provider::check_outer_fan_out_delegates([&self.gemm_capabilities])?;
754 return Ok(ParallelMode::Sequential);
755 }
756 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
757 return Ok(entry.provider_default_compatibility_mode());
758 }
759 self.validate_strict_capability(self.gemm_capabilities, entry.thread_budget().get())?;
760 entry.preferred_provider_mode(|mode| self.gemm_capabilities.accepts_mode(mode))
761 }
762
763 #[allow(clippy::too_many_arguments)]
764 fn execute_into(
765 &self,
766 bundle_identity: &Arc<CpuProviderBundleInner>,
767 entry: &CpuOperationEntry<'_>,
768 entered: Option<&CpuExecutionContext<'_>>,
769 buffers: &mut BufferPool,
770 cache: &mut GemmAnalysisCache,
771 cache_slot: Option<usize>,
772 lhs: TensorRead<'_>,
773 rhs: TensorRead<'_>,
774 config: &DotGeneralConfig,
775 accumulation: DotGeneralAccumulation,
776 output: TensorWrite<'_>,
777 ) -> Result<()> {
778 let validated = validate_dot_general(&lhs, &rhs, &output, config, accumulation)?;
779 let mode = self
780 .dot_general_mode(entry, entered)
781 .map_err(|error| Error::backend_source(OP, error))?;
782 let mode = self.batched_dot_mode(entry, entered, mode, dot_batch_items(&lhs, config)?)?;
783 cache.bind_provider_bundle(bundle_identity);
784 entry
785 .enter_or_reuse(entered, mode, |provider_context| {
786 self.execute_into_validated(
787 provider_context,
788 validated,
789 buffers,
790 cache,
791 cache_slot,
792 lhs,
793 rhs,
794 config,
795 accumulation,
796 output,
797 )
798 })
799 .map_err(|error| Error::backend_source(OP, error))?
800 }
801
802 #[allow(clippy::too_many_arguments)]
806 fn execute_into_validated(
807 &self,
808 provider_context: &CpuExecutionContext<'_>,
809 validated: ValidatedDotGeneral<'_>,
810 buffers: &mut BufferPool,
811 cache: &mut GemmAnalysisCache,
812 cache_slot: Option<usize>,
813 lhs: TensorRead<'_>,
814 rhs: TensorRead<'_>,
815 config: &DotGeneralConfig,
816 accumulation: DotGeneralAccumulation,
817 mut output: TensorWrite<'_>,
818 ) -> Result<()> {
819 if let Some(general) = &self.general {
820 let request = validated.request(&lhs, &rhs, &mut output, accumulation);
821 match general.dot_general(provider_context, request)? {
822 CpuProviderOutcome::Executed => return Ok(()),
823 CpuProviderOutcome::Unsupported(reason) => {
824 if self.general_policy == GeneralContractionPolicy::Required {
825 return Err(unsupported_provider_error(
826 "required general-contraction",
827 reason,
828 ));
829 }
830 }
831 }
832 }
833
834 if is_all_batch_contraction(&lhs, &rhs, config) {
838 return execute_all_batch_elementwise(
839 provider_context,
840 buffers,
841 &lhs,
842 &rhs,
843 config,
844 accumulation,
845 output,
846 );
847 }
848
849 if let Some(plan) =
850 crate::gemm::prepare_provider_gemm(cache, cache_slot, &lhs, &rhs, &output, config)?
851 {
852 match execute_gemm_plan(
853 self.gemm.as_ref(),
854 provider_context,
855 plan,
856 &lhs,
857 &rhs,
858 accumulation,
859 &mut output,
860 )? {
861 CpuProviderOutcome::Executed => return Ok(()),
862 CpuProviderOutcome::Unsupported(reason)
863 if !canonical_gemm_fallback_supported(reason) =>
864 {
865 return Err(unsupported_provider_error("GEMM", reason));
866 }
867 CpuProviderOutcome::Unsupported(_) => {}
868 }
869 }
870
871 self.execute_canonical_gemm(
872 provider_context,
873 buffers,
874 cache,
875 cache_slot,
876 &lhs,
877 &rhs,
878 config,
879 accumulation,
880 &mut output,
881 )
882 }
883
884 #[allow(clippy::too_many_arguments)]
887 fn with_canonical_operands<R>(
888 &self,
889 provider_context: &CpuExecutionContext<'_>,
890 buffers: &mut BufferPool,
891 lhs: &TensorRead<'_>,
892 rhs: &TensorRead<'_>,
893 config: &DotGeneralConfig,
894 accumulation: DotGeneralAccumulation,
895 execute: impl FnOnce(
896 &TensorRead<'_>,
897 &TensorRead<'_>,
898 &DotGeneralConfig,
899 DotGeneralAccumulation,
900 ) -> Result<R>,
901 ) -> Result<R> {
902 let (lhs_perm, rhs_perm, canonical_config) =
903 crate::gemm::canonical_gemm_layout(config, lhs.shape().len(), rhs.shape().len());
904 let lhs_canonical = self.canonical_operand(
905 provider_context,
906 buffers,
907 lhs,
908 &lhs_perm,
909 accumulation.lhs_conj,
910 )?;
911 let rhs_canonical = match self.canonical_operand(
912 provider_context,
913 buffers,
914 rhs,
915 &rhs_perm,
916 accumulation.rhs_conj,
917 ) {
918 Ok(operand) => operand,
919 Err(error) => {
920 lhs_canonical.reclaim(buffers);
921 return Err(error);
922 }
923 };
924 let result = execute(
925 &lhs_canonical.read(),
926 &rhs_canonical.read(),
927 &canonical_config,
928 DotGeneralAccumulation {
929 lhs_conj: false,
930 rhs_conj: false,
931 ..accumulation
932 },
933 );
934 lhs_canonical.reclaim(buffers);
935 rhs_canonical.reclaim(buffers);
936 result
937 }
938
939 fn canonical_operand<'input>(
942 &self,
943 provider_context: &CpuExecutionContext<'_>,
944 buffers: &mut BufferPool,
945 input: &TensorRead<'input>,
946 permutation: &[usize],
947 conjugate: bool,
948 ) -> Result<CanonicalOperand<'input>> {
949 if !conjugate {
950 let permuted = TensorRead::from_view(transposed_read_view(input, permutation)?);
951 if permuted.is_col_major_contiguous()? {
952 return Ok(CanonicalOperand::Borrowed(permuted));
953 }
954 }
955 materialize_canonical_operand(
956 self.layout.as_ref(),
957 provider_context,
958 buffers,
959 input,
960 permutation,
961 conjugate,
962 )
963 .map(CanonicalOperand::Packed)
964 }
965
966 #[allow(clippy::too_many_arguments)]
967 fn execute_canonical_gemm(
968 &self,
969 provider_context: &CpuExecutionContext<'_>,
970 buffers: &mut BufferPool,
971 cache: &mut GemmAnalysisCache,
972 cache_slot: Option<usize>,
973 lhs: &TensorRead<'_>,
974 rhs: &TensorRead<'_>,
975 config: &DotGeneralConfig,
976 accumulation: DotGeneralAccumulation,
977 output: &mut TensorWrite<'_>,
978 ) -> Result<()> {
979 self.with_canonical_operands(
980 provider_context,
981 buffers,
982 lhs,
983 rhs,
984 config,
985 accumulation,
986 |lhs, rhs, canonical_config, canonical_accumulation| {
987 let plan = match crate::gemm::canonical_provider_gemm_plan_into(
990 lhs,
991 rhs,
992 output,
993 canonical_config,
994 )? {
995 Some(plan) => Some(plan),
996 None => crate::gemm::prepare_provider_gemm_canonical(
997 cache,
998 cache_slot,
999 lhs,
1000 rhs,
1001 output,
1002 canonical_config,
1003 )?,
1004 };
1005 let Some(plan) = plan else {
1006 return Err(Error::unsupported(
1007 OP,
1008 "configured CPU layout-plus-GEMM path cannot represent the canonical contraction",
1009 ));
1010 };
1011 match execute_gemm_plan(
1012 self.gemm.as_ref(),
1013 provider_context,
1014 plan,
1015 lhs,
1016 rhs,
1017 canonical_accumulation,
1018 output,
1019 )? {
1020 CpuProviderOutcome::Executed => Ok(()),
1021 CpuProviderOutcome::Unsupported(reason) => {
1022 Err(unsupported_provider_error("GEMM", reason))
1023 }
1024 }
1025 },
1026 )
1027 }
1028
1029 #[allow(clippy::too_many_arguments)]
1044 pub(crate) fn execute_dot_into_uninit(
1045 &self,
1046 bundle_identity: &Arc<CpuProviderBundleInner>,
1047 entry: &CpuOperationEntry<'_>,
1048 entered: Option<&CpuExecutionContext<'_>>,
1049 buffers: &mut BufferPool,
1050 cache: &mut GemmAnalysisCache,
1051 cache_slot: Option<usize>,
1052 lhs: &TensorRead<'_>,
1053 rhs: &TensorRead<'_>,
1054 config: &DotGeneralConfig,
1055 accumulation: DotGeneralAccumulation,
1056 output_shape: &[usize],
1057 output_bytes: &mut [MaybeUninit<u8>],
1058 ) -> Result<CpuProviderOutcome> {
1059 let Some(witness) = self.gemm.uninit_provider() else {
1060 return Err(Error::unsupported(
1061 OP,
1062 "configured CPU GEMM provider does not expose the uninitialized-output contract",
1063 ));
1064 };
1065 let mode = self
1066 .dot_general_mode(entry, entered)
1067 .map_err(|error| Error::backend_source(OP, error))?;
1068 let items = dot_batch_items(lhs, config)?;
1069 let mode = self.batched_dot_mode(entry, entered, mode, items)?;
1070 if items > 1
1073 && effective_batch_policy(entry, entered).strategy()
1074 == CpuBatchStrategy::WholeBatchVendor
1075 {
1076 return Ok(CpuProviderOutcome::Unsupported(
1077 CpuProviderUnsupported::StridedBatch,
1078 ));
1079 }
1080 if is_all_batch_contraction(lhs, rhs, config) {
1081 return Ok(CpuProviderOutcome::Unsupported(
1082 CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Output),
1083 ));
1084 }
1085 cache.bind_provider_bundle(bundle_identity);
1086 entry
1087 .enter_or_reuse(entered, mode, |provider_context| {
1088 if let Some(plan) = crate::gemm::prepare_provider_gemm_into_uninit(
1089 cache,
1090 cache_slot,
1091 lhs,
1092 rhs,
1093 output_shape,
1094 config,
1095 )? {
1096 match execute_gemm_plan_into_uninit(
1097 witness,
1098 provider_context,
1099 plan,
1100 lhs,
1101 rhs,
1102 accumulation,
1103 &mut *output_bytes,
1104 )? {
1105 CpuProviderOutcome::Executed => return Ok(CpuProviderOutcome::Executed),
1106 CpuProviderOutcome::Unsupported(reason)
1107 if !canonical_gemm_fallback_supported(reason) =>
1108 {
1109 return Ok(CpuProviderOutcome::Unsupported(reason));
1110 }
1111 CpuProviderOutcome::Unsupported(_) => {}
1113 }
1114 }
1115 self.with_canonical_operands(
1116 provider_context,
1117 buffers,
1118 lhs,
1119 rhs,
1120 config,
1121 accumulation,
1122 |lhs, rhs, canonical_config, canonical_accumulation| {
1123 let plan = match crate::gemm::canonical_provider_gemm_plan_uninit(
1124 lhs,
1125 rhs,
1126 output_shape,
1127 canonical_config,
1128 )? {
1129 Some(plan) => Some(plan),
1130 None => crate::gemm::prepare_provider_gemm_canonical_into_uninit(
1131 cache,
1132 cache_slot,
1133 lhs,
1134 rhs,
1135 output_shape,
1136 canonical_config,
1137 )?,
1138 };
1139 let Some(plan) = plan else {
1140 return Ok(CpuProviderOutcome::Unsupported(
1141 CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Output),
1142 ));
1143 };
1144 execute_gemm_plan_into_uninit(
1145 witness,
1146 provider_context,
1147 plan,
1148 lhs,
1149 rhs,
1150 canonical_accumulation,
1151 output_bytes,
1152 )
1153 },
1154 )
1155 })
1156 .map_err(|error| Error::backend_source(OP, error))?
1157 }
1158
1159 #[allow(clippy::redundant_closure)]
1160 fn execute_grouped(
1161 &self,
1162 entry: &CpuOperationEntry<'_>,
1163 entered: Option<&CpuExecutionContext<'_>>,
1164 lhs: TensorRead<'_>,
1165 rhs: TensorRead<'_>,
1166 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
1167 mut output: TensorWrite<'_>,
1168 ) -> Result<()> {
1169 tenferro_tensor::backend::validate_grouped_gemm(
1170 &lhs,
1171 &rhs,
1172 &output,
1173 config,
1174 "grouped_gemm",
1175 )?;
1176 let policy = effective_batch_policy(entry, entered);
1177 let jobs = config.jobs().len();
1178 let strategy = if jobs > 1 {
1180 policy.strategy()
1181 } else {
1182 CpuBatchStrategy::Auto
1183 };
1184 let fan_out = match strategy {
1185 CpuBatchStrategy::Auto
1186 if self.grouped_scheduling == GroupedGemmScheduling::EngineOuter =>
1187 {
1188 match entered {
1189 None => (entry.supports_outer()
1190 && policy
1191 .thresholds()
1192 .fans_out(jobs, jobs.min(entry.thread_budget().get())))
1193 .then_some(crate::provider::CpuOuterFanOut::Executor(*entry)),
1194 Some(context) if context.can_fan_out_lanes() => {
1197 auto_grouped_lane_count(
1198 policy.thresholds(),
1199 config.jobs(),
1200 context.thread_budget().get(),
1201 )
1202 .map(|_| crate::provider::CpuOuterFanOut::Lanes(*context))
1203 }
1204 Some(_) => None,
1205 }
1206 }
1207 CpuBatchStrategy::OuterParallel => Some(match entered {
1208 None if entry.supports_outer() => crate::provider::CpuOuterFanOut::Executor(*entry),
1209 Some(context) if context.can_fan_out_lanes() => {
1210 crate::provider::CpuOuterFanOut::Lanes(*context)
1211 }
1212 _ => {
1213 return Err(strategy_unavailable(
1214 "grouped_gemm",
1215 strategy,
1216 "the selected CPU domain cannot fan out (one thread, or no outer-capable executor)",
1217 ))
1218 }
1219 }),
1220 _ => None,
1221 };
1222 if let Some(fan_out) = fan_out {
1223 let checked = crate::provider::check_outer_fan_out_delegates([&self.gemm_capabilities])
1225 .map_err(|error| Error::backend_source("grouped_gemm", error))?;
1226 let chunks = match fan_out {
1234 crate::provider::CpuOuterFanOut::Lanes(context) => auto_grouped_lane_count(
1235 policy.thresholds(),
1236 config.jobs(),
1237 context.thread_budget().get(),
1238 )
1239 .unwrap_or_else(|| jobs.min(context.thread_budget().get())),
1240 crate::provider::CpuOuterFanOut::Executor(_) => jobs,
1241 };
1242 macro_rules! outer_typed {
1243 ($variant:ident, $storage:expr, $base:expr) => {
1244 execute_grouped_outer_typed(
1245 self.gemm.as_ref(),
1246 checked,
1247 fan_out,
1248 chunks,
1249 &lhs,
1250 &rhs,
1251 config,
1252 $storage,
1253 $base,
1254 |view| TensorViewMut::$variant(view),
1255 )
1256 };
1257 }
1258
1259 return match &mut output {
1260 TensorWrite::Tensor(tensor) => match tensor.dtype() {
1261 DType::F32 => {
1262 outer_typed!(F32, dot_write_operand::<f32>(tensor)?.host_data_mut()?, 0)
1263 }
1264 DType::F64 => {
1265 outer_typed!(F64, dot_write_operand::<f64>(tensor)?.host_data_mut()?, 0)
1266 }
1267 DType::C32 => outer_typed!(
1268 C32,
1269 dot_write_operand::<Complex32>(tensor)?.host_data_mut()?,
1270 0
1271 ),
1272 DType::C64 => outer_typed!(
1273 C64,
1274 dot_write_operand::<Complex64>(tensor)?.host_data_mut()?,
1275 0
1276 ),
1277 _ => Err(unsupported_provider_error(
1278 "grouped-GEMM",
1279 CpuProviderUnsupported::DType(tensor.dtype()),
1280 )),
1281 },
1282 TensorWrite::View(TensorViewMut::F32(output)) => {
1283 let base = output.offset();
1284 outer_typed!(F32, output.host_storage_mut()?, base)
1285 }
1286 TensorWrite::View(TensorViewMut::F64(output)) => {
1287 let base = output.offset();
1288 outer_typed!(F64, output.host_storage_mut()?, base)
1289 }
1290 TensorWrite::View(TensorViewMut::C32(output)) => {
1291 let base = output.offset();
1292 outer_typed!(C32, output.host_storage_mut()?, base)
1293 }
1294 TensorWrite::View(TensorViewMut::C64(output)) => {
1295 let base = output.offset();
1296 outer_typed!(C64, output.host_storage_mut()?, base)
1297 }
1298 _ => Err(unsupported_provider_error(
1299 "grouped-GEMM",
1300 CpuProviderUnsupported::DType(output.dtype()),
1301 )),
1302 };
1303 }
1304 let mut mode = self
1305 .grouped_mode(entry, entered)
1306 .map_err(|error| Error::backend_source("grouped_gemm", error))?;
1307 if strategy == CpuBatchStrategy::Sequential {
1308 if !self
1309 .gemm_capabilities
1310 .accepts_mode(ParallelMode::Sequential)
1311 {
1312 return Err(strategy_unavailable(
1313 "grouped_gemm",
1314 strategy,
1315 "the GEMM provider does not accept sequential calls",
1316 ));
1317 }
1318 mode = ParallelMode::Sequential;
1319 }
1320 let vendor_batch = vendor_batch_for(policy, strategy);
1321 entry
1322 .enter_or_reuse(entered, mode, |provider_context| {
1323 let request = CpuGroupedGemmRequest::new(
1324 &lhs,
1325 &rhs,
1326 &mut output,
1327 config.jobs(),
1328 config.accumulation(),
1329 )
1330 .with_vendor_batch(vendor_batch);
1331 match self.gemm.grouped_gemm(provider_context, request)? {
1332 CpuProviderOutcome::Executed => Ok(()),
1333 CpuProviderOutcome::Unsupported(reason) => {
1334 Err(unsupported_provider_error("grouped-GEMM", reason))
1335 }
1336 }
1337 })
1338 .map_err(|error| Error::backend_source("grouped_gemm", error))?
1339 }
1340}
1341
1342fn execute_gemm_plan(
1343 provider: &dyn CpuGemmProvider,
1344 context: &CpuExecutionContext<'_>,
1345 plan: crate::gemm::ProviderGemmPlan,
1346 lhs: &TensorRead<'_>,
1347 rhs: &TensorRead<'_>,
1348 accumulation: DotGeneralAccumulation,
1349 output: &mut TensorWrite<'_>,
1350) -> Result<CpuProviderOutcome> {
1351 let batch_count = plan.batch_count();
1352 let policy = context.batch_policy();
1353 let strategy = if batch_count > 1 {
1354 policy.strategy()
1355 } else {
1356 CpuBatchStrategy::Auto
1357 };
1358 let vendor_batch = match strategy {
1363 CpuBatchStrategy::Auto if batch_count > 1 => crate::provider::CpuVendorBatch::Forbidden,
1364 _ => vendor_batch_for(policy, strategy),
1365 };
1366 if batch_count > 1
1367 && matches!(
1368 strategy,
1369 CpuBatchStrategy::Auto | CpuBatchStrategy::OuterParallel
1370 )
1371 {
1372 if let Some(outcome) =
1373 try_execute_gemm_plan_on_lanes(provider, context, plan, lhs, rhs, accumulation, output)?
1374 {
1375 return Ok(outcome);
1376 }
1377 if strategy == CpuBatchStrategy::OuterParallel {
1378 return Err(forced_lanes_unavailable());
1379 }
1380 }
1381 let request = plan
1382 .request(lhs, rhs, output, accumulation)
1383 .with_vendor_batch(vendor_batch);
1384 let outcome = if batch_count == 1 {
1385 provider.gemm(context, request)?
1386 } else {
1387 provider.strided_batched_gemm(context, request)?
1388 };
1389 Ok(outcome)
1390}
1391
1392fn forced_lanes_unavailable() -> Error {
1401 strategy_unavailable(
1402 OP,
1403 CpuBatchStrategy::OuterParallel,
1404 "this strided batch cannot be split over outer lanes (one thread or a nested lane, \
1405 overlapping or reversed output items, or a provider that runs its own threads)",
1406 )
1407}
1408
1409fn strided_batch_lanes(
1410 context: &CpuExecutionContext<'_>,
1411 plan: crate::gemm::ProviderGemmPlan,
1412) -> Option<usize> {
1413 let batch = plan.batch_count();
1414 if batch <= 1 || !context.can_fan_out_lanes() {
1415 return None;
1416 }
1417 let threads = context.thread_budget().get();
1418 let policy = context.batch_policy();
1419 match policy.strategy() {
1420 CpuBatchStrategy::Auto => {
1421 let thresholds = policy.thresholds();
1422 let item_ns = thresholds.lane_item_ns(plan.rows(), plan.columns(), plan.contracted());
1423 thresholds.auto_lanes(batch, item_ns.saturating_mul(batch), threads)
1424 }
1425 CpuBatchStrategy::OuterParallel => Some(threads.min(batch)).filter(|&lanes| lanes >= 2),
1426 _ => None,
1427 }
1428}
1429
1430fn auto_grouped_lane_count(
1435 thresholds: crate::CpuBatchThresholds,
1436 jobs: &[tenferro_tensor::backend::GroupedGemmJob],
1437 threads: usize,
1438) -> Option<usize> {
1439 let total_ns = jobs.iter().fold(0usize, |total, job| {
1440 total.saturating_add(thresholds.lane_item_ns(job.rows(), job.cols(), job.contracted()))
1441 });
1442 thresholds.auto_lanes(jobs.len(), total_ns, threads)
1443}
1444
1445fn try_execute_gemm_plan_on_lanes(
1451 provider: &dyn CpuGemmProvider,
1452 context: &CpuExecutionContext<'_>,
1453 plan: crate::gemm::ProviderGemmPlan,
1454 lhs: &TensorRead<'_>,
1455 rhs: &TensorRead<'_>,
1456 accumulation: DotGeneralAccumulation,
1457 output: &mut TensorWrite<'_>,
1458) -> Result<Option<CpuProviderOutcome>> {
1459 let Some(lanes) = strided_batch_lanes(context, plan) else {
1460 return Ok(None);
1461 };
1462 if crate::provider::check_outer_fan_out_delegates([&provider.execution_capabilities()]).is_err()
1463 {
1464 return Ok(None);
1465 }
1466 let Some(item_span) = output_item_span(plan) else {
1467 return Ok(None);
1468 };
1469 macro_rules! typed {
1470 ($ty:ty, $variant:ident, $storage:expr) => {
1471 execute_gemm_chunks_on_lanes::<$ty>(
1472 provider,
1473 context,
1474 plan,
1475 lhs,
1476 rhs,
1477 accumulation,
1478 $storage,
1479 item_span,
1480 lanes,
1481 |view| TensorViewMut::$variant(view),
1482 )
1483 };
1484 }
1485 match output {
1486 TensorWrite::Tensor(tensor) => match tensor.dtype() {
1487 DType::F32 => typed!(f32, F32, dot_write_operand::<f32>(tensor)?.host_data_mut()?),
1488 DType::F64 => typed!(f64, F64, dot_write_operand::<f64>(tensor)?.host_data_mut()?),
1489 DType::C32 => typed!(
1490 Complex32,
1491 C32,
1492 dot_write_operand::<Complex32>(tensor)?.host_data_mut()?
1493 ),
1494 DType::C64 => typed!(
1495 Complex64,
1496 C64,
1497 dot_write_operand::<Complex64>(tensor)?.host_data_mut()?
1498 ),
1499 _ => Ok(None),
1500 },
1501 TensorWrite::View(TensorViewMut::F32(view)) => typed!(f32, F32, view.host_storage_mut()?),
1502 TensorWrite::View(TensorViewMut::F64(view)) => typed!(f64, F64, view.host_storage_mut()?),
1503 TensorWrite::View(TensorViewMut::C32(view)) => {
1504 typed!(Complex32, C32, view.host_storage_mut()?)
1505 }
1506 TensorWrite::View(TensorViewMut::C64(view)) => {
1507 typed!(Complex64, C64, view.host_storage_mut()?)
1508 }
1509 TensorWrite::View(_) => Ok(None),
1510 }
1511}
1512
1513fn output_item_span(plan: crate::gemm::ProviderGemmPlan) -> Option<usize> {
1516 let layout = plan.output_layout();
1517 let positive = |stride: isize| usize::try_from(stride).ok().filter(|&stride| stride > 0);
1518 let (row, column, batch) = (
1519 positive(layout.row_stride())?,
1520 positive(layout.column_stride())?,
1521 positive(layout.batch_stride())?,
1522 );
1523 let span = plan
1524 .rows()
1525 .checked_sub(1)?
1526 .checked_mul(row)?
1527 .checked_add(plan.columns().checked_sub(1)?.checked_mul(column)?)?
1528 .checked_add(1)?;
1529 (span <= batch).then_some(span)
1530}
1531
1532#[allow(clippy::too_many_arguments)]
1533fn execute_gemm_chunks_on_lanes<T>(
1534 provider: &dyn CpuGemmProvider,
1535 context: &CpuExecutionContext<'_>,
1536 plan: crate::gemm::ProviderGemmPlan,
1537 lhs: &TensorRead<'_>,
1538 rhs: &TensorRead<'_>,
1539 accumulation: DotGeneralAccumulation,
1540 storage: &mut [T],
1541 item_span: usize,
1542 lanes: usize,
1543 wrap: for<'a> fn(tenferro_tensor::TypedTensorViewMut<'a, T>) -> TensorViewMut<'a>,
1544) -> Result<Option<CpuProviderOutcome>>
1545where
1546 T: Send + Sync + 'static,
1547{
1548 let batch = plan.batch_count();
1549 let layout = plan.output_layout();
1550 let (Ok(first), Ok(batch_stride)) = (
1551 usize::try_from(layout.offset()),
1552 usize::try_from(layout.batch_stride()),
1553 ) else {
1554 return Ok(None);
1555 };
1556 let mut chunks = Vec::with_capacity(lanes);
1558 let mut rest = storage;
1559 let mut cursor = 0usize;
1560 let mut start = 0usize;
1561 for lane in 0..lanes {
1562 let len = batch / lanes + usize::from(lane < batch % lanes);
1563 let (Some(begin), Some(end)) = (
1564 start
1565 .checked_mul(batch_stride)
1566 .and_then(|value| value.checked_add(first)),
1567 (start + len - 1)
1568 .checked_mul(batch_stride)
1569 .and_then(|value| value.checked_add(first))
1570 .and_then(|value| value.checked_add(item_span)),
1571 ) else {
1572 return Ok(None);
1573 };
1574 if end - cursor > rest.len() {
1575 return Ok(None);
1576 }
1577 let (_, tail) = std::mem::take(&mut rest).split_at_mut(begin - cursor);
1578 let (chunk, tail) = tail.split_at_mut(end - begin);
1579 rest = tail;
1580 cursor = end;
1581 let Some(chunk_plan) = plan.batch_chunk(start, len, 0) else {
1582 return Ok(None);
1583 };
1584 let shape = [plan.rows(), plan.columns(), len];
1585 let strides = [
1586 layout.row_stride(),
1587 layout.column_stride(),
1588 layout.batch_stride(),
1589 ];
1590 let view = tenferro_tensor::TypedTensorViewMut::from_slice(shape, strides, 0, chunk)?;
1591 chunks.push((chunk_plan, view));
1592 start += len;
1593 }
1594
1595 let outcomes = std::sync::Mutex::new(Vec::with_capacity(lanes));
1596 context.with_outer_lanes(chunks, |(chunk_plan, view), lane| {
1597 let mut chunk_output = TensorWrite::from_view(wrap(view));
1598 let request = chunk_plan
1599 .request(lhs, rhs, &mut chunk_output, accumulation)
1600 .with_vendor_batch(crate::provider::CpuVendorBatch::Forbidden);
1601 let outcome = provider.strided_batched_gemm(lane, request);
1602 outcomes
1603 .lock()
1604 .unwrap_or_else(std::sync::PoisonError::into_inner)
1605 .push(outcome);
1606 });
1607 let outcomes = outcomes
1608 .into_inner()
1609 .unwrap_or_else(std::sync::PoisonError::into_inner);
1610 let mut unsupported = None;
1611 for outcome in outcomes {
1612 match outcome? {
1613 CpuProviderOutcome::Executed => {}
1614 CpuProviderOutcome::Unsupported(reason) => unsupported = Some(reason),
1615 }
1616 }
1617 match unsupported {
1618 None => Ok(Some(CpuProviderOutcome::Executed)),
1619 Some(reason) if accumulation_is_overwrite(accumulation)? => {
1623 Ok(Some(CpuProviderOutcome::Unsupported(reason)))
1624 }
1625 Some(reason) => Err(unsupported_provider_error("GEMM", reason)),
1626 }
1627}
1628
1629fn accumulation_is_overwrite(accumulation: DotGeneralAccumulation) -> Result<bool> {
1630 let overwrite = DotGeneralAccumulation::overwrite(accumulation.alpha.dtype())?;
1631 Ok(accumulation.alpha == overwrite.alpha && accumulation.beta == overwrite.beta)
1632}
1633
1634fn execute_gemm_plan_into_uninit(
1635 witness: &dyn CpuUninitGemmProvider,
1636 context: &CpuExecutionContext<'_>,
1637 plan: crate::gemm::ProviderGemmPlan,
1638 lhs: &TensorRead<'_>,
1639 rhs: &TensorRead<'_>,
1640 accumulation: DotGeneralAccumulation,
1641 output_bytes: &mut [MaybeUninit<u8>],
1642) -> Result<CpuProviderOutcome> {
1643 let strategy = context.batch_policy().strategy();
1646 if plan.batch_count() > 1
1647 && matches!(
1648 strategy,
1649 CpuBatchStrategy::Auto | CpuBatchStrategy::OuterParallel
1650 )
1651 {
1652 if let Some(outcome) = try_execute_gemm_plan_into_uninit_on_lanes(
1653 witness,
1654 context,
1655 plan,
1656 lhs,
1657 rhs,
1658 accumulation,
1659 output_bytes,
1660 )? {
1661 return Ok(outcome);
1662 }
1663 if strategy == CpuBatchStrategy::OuterParallel {
1664 return Err(forced_lanes_unavailable());
1665 }
1666 }
1667 let request = plan.uninit_request(lhs, rhs, accumulation);
1668 unsafe { witness.gemm_into_uninit(context, request, output_bytes) }
1673}
1674
1675fn try_execute_gemm_plan_into_uninit_on_lanes(
1680 witness: &dyn CpuUninitGemmProvider,
1681 context: &CpuExecutionContext<'_>,
1682 plan: crate::gemm::ProviderGemmPlan,
1683 lhs: &TensorRead<'_>,
1684 rhs: &TensorRead<'_>,
1685 accumulation: DotGeneralAccumulation,
1686 output_bytes: &mut [MaybeUninit<u8>],
1687) -> Result<Option<CpuProviderOutcome>> {
1688 let batch = plan.batch_count();
1689 let Some(lanes) = strided_batch_lanes(context, plan) else {
1690 return Ok(None);
1691 };
1692 if crate::provider::check_outer_fan_out_delegates([&witness.execution_capabilities()]).is_err()
1693 {
1694 return Ok(None);
1695 }
1696 let element_size = match lhs.dtype() {
1697 DType::F32 => std::mem::size_of::<f32>(),
1698 DType::F64 => std::mem::size_of::<f64>(),
1699 DType::C32 => std::mem::size_of::<Complex32>(),
1700 DType::C64 => std::mem::size_of::<Complex64>(),
1701 _ => return Ok(None),
1702 };
1703 let Some(item_span) = output_item_span(plan) else {
1704 return Ok(None);
1705 };
1706 let layout = plan.output_layout();
1707 let (Ok(first), Ok(batch_stride)) = (
1708 usize::try_from(layout.offset()),
1709 usize::try_from(layout.batch_stride()),
1710 ) else {
1711 return Ok(None);
1712 };
1713 let mut chunks = Vec::with_capacity(lanes);
1715 let mut rest = output_bytes;
1716 let mut cursor = 0usize;
1717 let mut start = 0usize;
1718 for lane in 0..lanes {
1719 let len = batch / lanes + usize::from(lane < batch % lanes);
1720 let (Some(begin), Some(end)) = (
1721 start
1722 .checked_mul(batch_stride)
1723 .and_then(|value| value.checked_add(first))
1724 .and_then(|value| value.checked_mul(element_size)),
1725 (start + len - 1)
1726 .checked_mul(batch_stride)
1727 .and_then(|value| value.checked_add(first))
1728 .and_then(|value| value.checked_add(item_span))
1729 .and_then(|value| value.checked_mul(element_size)),
1730 ) else {
1731 return Ok(None);
1732 };
1733 if end - cursor > rest.len() {
1734 return Ok(None);
1735 }
1736 let (_, tail) = std::mem::take(&mut rest).split_at_mut(begin - cursor);
1737 let (chunk, tail) = tail.split_at_mut(end - begin);
1738 rest = tail;
1739 cursor = end;
1740 let Some(chunk_plan) = plan.batch_chunk(start, len, 0) else {
1741 return Ok(None);
1742 };
1743 chunks.push((chunk_plan, chunk));
1744 start += len;
1745 }
1746
1747 let outcomes = std::sync::Mutex::new(Vec::with_capacity(lanes));
1748 context.with_outer_lanes(chunks, |(chunk_plan, chunk), lane| {
1749 let request = chunk_plan.uninit_request(lhs, rhs, accumulation);
1750 let outcome = unsafe { witness.gemm_into_uninit(lane, request, chunk) };
1754 outcomes
1755 .lock()
1756 .unwrap_or_else(std::sync::PoisonError::into_inner)
1757 .push(outcome);
1758 });
1759 let outcomes = outcomes
1760 .into_inner()
1761 .unwrap_or_else(std::sync::PoisonError::into_inner);
1762 let mut unsupported = None;
1763 for outcome in outcomes {
1764 match outcome? {
1765 CpuProviderOutcome::Executed => {}
1766 CpuProviderOutcome::Unsupported(reason) => unsupported = Some(reason),
1767 }
1768 }
1769 Ok(Some(match unsupported {
1772 None => CpuProviderOutcome::Executed,
1773 Some(reason) => CpuProviderOutcome::Unsupported(reason),
1774 }))
1775}
1776
1777fn canonical_gemm_fallback_supported(reason: CpuProviderUnsupported) -> bool {
1778 matches!(
1779 reason,
1780 CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Lhs)
1781 | CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Rhs)
1782 | CpuProviderUnsupported::Conjugation
1783 )
1784}
1785
1786fn transposed_read_view<'input>(
1787 input: &TensorRead<'input>,
1788 permutation: &[usize],
1789) -> Result<TensorView<'input>> {
1790 Ok(match input.clone().tensor_view() {
1791 TensorView::F32(view) => TensorView::F32(view.transpose_view(permutation)?),
1792 TensorView::F64(view) => TensorView::F64(view.transpose_view(permutation)?),
1793 TensorView::I32(view) => TensorView::I32(view.transpose_view(permutation)?),
1794 TensorView::I64(view) => TensorView::I64(view.transpose_view(permutation)?),
1795 TensorView::Bool(view) => TensorView::Bool(view.transpose_view(permutation)?),
1796 TensorView::C32(view) => TensorView::C32(view.transpose_view(permutation)?),
1797 TensorView::C64(view) => TensorView::C64(view.transpose_view(permutation)?),
1798 })
1799}
1800
1801fn pooled_zero_tensor<T>(buffers: &mut BufferPool, shape: Vec<usize>) -> Result<TypedTensor<T>>
1802where
1803 T: PoolScalar + Clone + 'static,
1804{
1805 let element_count =
1806 tenferro_tensor::validate::checked_shape_product(OP, "canonical operand", &shape)?;
1807 TypedTensor::from_vec_col_major(shape, T::pool_acquire_zeroed(buffers, element_count))
1808}
1809
1810fn allocate_canonical_operand(
1811 buffers: &mut BufferPool,
1812 dtype: DType,
1813 shape: Vec<usize>,
1814) -> Result<Tensor> {
1815 match dtype {
1816 DType::F32 => pooled_zero_tensor(buffers, shape).map(Tensor::from_typed::<f32>),
1817 DType::F64 => pooled_zero_tensor(buffers, shape).map(Tensor::from_typed::<f64>),
1818 DType::C32 => {
1819 pooled_zero_tensor(buffers, shape).map(Tensor::from_typed::<num_complex::Complex32>)
1820 }
1821 DType::C64 => {
1822 pooled_zero_tensor(buffers, shape).map(Tensor::from_typed::<num_complex::Complex64>)
1823 }
1824 dtype => Err(Error::unsupported_dtype(
1825 OP,
1826 dtype,
1827 crate::cpu_contraction_unsupported_dtype_message(dtype),
1828 )),
1829 }
1830}
1831
1832fn dot_write_operand<T: tenferro_tensor::TensorScalar>(
1834 tensor: &mut Tensor,
1835) -> crate::Result<&mut TypedTensor<T>> {
1836 let dtype = tensor.dtype();
1837 tensor.as_typed_mut::<T>().ok_or_else(|| {
1838 unsupported_provider_error("grouped-GEMM", CpuProviderUnsupported::DType(dtype))
1839 })
1840}
1841
1842fn validated_operand<'a, T: tenferro_tensor::TensorScalar>(
1844 tensor: &'a Tensor,
1845 op: &'static str,
1846 message: &'static str,
1847) -> crate::Result<&'a TypedTensor<T>> {
1848 tensor
1849 .as_typed::<T>()
1850 .ok_or_else(|| crate::Error::unsupported_dtype(op, tensor.dtype(), message))
1851}
1852
1853fn materialize_canonical_operand(
1854 provider: &dyn CpuLayoutTransformProvider,
1855 context: &CpuExecutionContext<'_>,
1856 buffers: &mut BufferPool,
1857 input: &TensorRead<'_>,
1858 permutation: &[usize],
1859 conjugate: bool,
1860) -> Result<Tensor> {
1861 let input_view = transposed_read_view(input, permutation)?;
1862 let dtype = input_view.dtype();
1863 let input = TensorRead::from_view(input_view);
1864 if let Some(witness) = provider.uninit_provider() {
1865 let mut output = UninitTensor::acquire(buffers, dtype, input.shape().to_vec())?;
1866 let outcome = {
1867 let output_bytes = output.as_uninit_bytes_mut();
1868 unsafe {
1873 witness.materialize_into_uninit(
1874 context,
1875 &input,
1876 CpuLayoutTransformIntent::CanonicalColumnMajor,
1877 conjugate,
1878 output_bytes,
1879 )
1880 }
1881 };
1882 match outcome {
1883 Ok(CpuProviderOutcome::Executed) => {
1884 return unsafe { output.assume_init() };
1887 }
1888 Ok(CpuProviderOutcome::Unsupported(_)) => {
1889 }
1892 Err(error) => return Err(error),
1893 }
1894 }
1895 let shape = input.shape().to_vec();
1896 materialize_canonical_operand_zeroed(provider, context, buffers, &input, shape, conjugate)
1897}
1898
1899fn materialize_canonical_operand_zeroed(
1900 provider: &dyn CpuLayoutTransformProvider,
1901 context: &CpuExecutionContext<'_>,
1902 buffers: &mut BufferPool,
1903 input: &TensorRead<'_>,
1904 shape: Vec<usize>,
1905 conjugate: bool,
1906) -> Result<Tensor> {
1907 let mut output = allocate_canonical_operand(buffers, input.dtype(), shape)?;
1908 let outcome = {
1909 let mut output_write = TensorWrite::from_tensor(&mut output);
1910 let request = CpuLayoutTransformRequest::new(
1911 input,
1912 &mut output_write,
1913 CpuLayoutTransformIntent::CanonicalColumnMajor,
1914 conjugate,
1915 );
1916 provider.materialize(context, request)
1917 };
1918 match outcome {
1919 Ok(CpuProviderOutcome::Executed) => Ok(output),
1920 Ok(CpuProviderOutcome::Unsupported(reason)) => {
1921 crate::backend::reclaim_tensor(buffers, output);
1922 Err(unsupported_provider_error("layout-transform", reason))
1923 }
1924 Err(error) => {
1925 crate::backend::reclaim_tensor(buffers, output);
1926 Err(error)
1927 }
1928 }
1929}
1930
1931pub(crate) enum UninitTensor {
1938 F32(PooledUninitOutput<f32>),
1939 F64(PooledUninitOutput<f64>),
1940 C32(PooledUninitOutput<Complex32>),
1941 C64(PooledUninitOutput<Complex64>),
1942}
1943
1944impl UninitTensor {
1945 pub(crate) fn acquire(buffers: &BufferPool, dtype: DType, shape: Vec<usize>) -> Result<Self> {
1946 match dtype {
1947 DType::F32 => Ok(Self::F32(PooledUninitOutput::new(buffers, shape)?)),
1948 DType::F64 => Ok(Self::F64(PooledUninitOutput::new(buffers, shape)?)),
1949 DType::C32 => Ok(Self::C32(PooledUninitOutput::new(buffers, shape)?)),
1950 DType::C64 => Ok(Self::C64(PooledUninitOutput::new(buffers, shape)?)),
1951 dtype => Err(Error::unsupported_dtype(
1952 OP,
1953 dtype,
1954 crate::cpu_contraction_unsupported_dtype_message(dtype),
1955 )),
1956 }
1957 }
1958
1959 pub(crate) fn as_uninit_bytes_mut(&mut self) -> &mut [MaybeUninit<u8>] {
1960 match self {
1961 Self::F32(output) => output.as_uninit_bytes_mut(),
1962 Self::F64(output) => output.as_uninit_bytes_mut(),
1963 Self::C32(output) => output.as_uninit_bytes_mut(),
1964 Self::C64(output) => output.as_uninit_bytes_mut(),
1965 }
1966 }
1967
1968 pub(crate) unsafe fn assume_init(self) -> Result<Tensor> {
1974 unsafe {
1977 match self {
1978 Self::F32(output) => output.assume_init().map(Tensor::from_typed::<f32>),
1979 Self::F64(output) => output.assume_init().map(Tensor::from_typed::<f64>),
1980 Self::C32(output) => output
1981 .assume_init()
1982 .map(Tensor::from_typed::<num_complex::Complex32>),
1983 Self::C64(output) => output
1984 .assume_init()
1985 .map(Tensor::from_typed::<num_complex::Complex64>),
1986 }
1987 }
1988 }
1989}
1990
1991fn checked_grouped_output_range(
1992 output_base: usize,
1993 output_len: usize,
1994 job: &tenferro_tensor::backend::GroupedGemmJob,
1995) -> Result<std::ops::Range<usize>> {
1996 let len = job.rows().checked_mul(job.cols()).ok_or_else(|| {
1997 Error::invalid_argument(
1998 "grouped_gemm",
1999 "jobs",
2000 "grouped-GEMM output span overflows usize",
2001 )
2002 })?;
2003 let start = output_base.checked_add(job.out_offset()).ok_or_else(|| {
2004 Error::invalid_argument(
2005 "grouped_gemm",
2006 "jobs",
2007 "grouped-GEMM output offset overflows usize",
2008 )
2009 })?;
2010 let end = start.checked_add(len).ok_or_else(|| {
2011 Error::invalid_argument(
2012 "grouped_gemm",
2013 "jobs",
2014 "grouped-GEMM output end overflows usize",
2015 )
2016 })?;
2017 if end > output_len {
2018 return Err(Error::invalid_argument(
2019 "grouped_gemm",
2020 "jobs",
2021 "grouped-GEMM output range exceeds host storage",
2022 ));
2023 }
2024 Ok(start..end)
2025}
2026
2027fn grouped_output_starts_increase(
2035 jobs: &[tenferro_tensor::backend::GroupedGemmJob],
2036 output_base: usize,
2037 output_len: usize,
2038) -> Result<bool> {
2039 let mut previous = None;
2040 let mut increasing = true;
2041 for job in jobs {
2042 let range = checked_grouped_output_range(output_base, output_len, job)?;
2043 if !range.is_empty() {
2044 increasing &= previous.is_none_or(|start| start < range.start);
2045 previous = Some(range.start);
2046 }
2047 }
2048 Ok(increasing)
2049}
2050
2051fn grouped_chunk(
2054 jobs: &[tenferro_tensor::backend::GroupedGemmJob],
2055 output_base: usize,
2056 output_len: usize,
2057) -> Result<(
2058 std::ops::Range<usize>,
2059 SmallVec<[tenferro_tensor::backend::GroupedGemmJob; 1]>,
2060)> {
2061 let mut union: Option<std::ops::Range<usize>> = None;
2062 for job in jobs {
2063 let range = checked_grouped_output_range(output_base, output_len, job)?;
2064 if !range.is_empty() {
2065 union = Some(union.map_or(range.clone(), |union| {
2066 union.start.min(range.start)..union.end.max(range.end)
2067 }));
2068 }
2069 }
2070 let union = union.unwrap_or(0..0);
2071 let rebased = jobs
2072 .iter()
2073 .map(|job| {
2074 let start = output_base + job.out_offset();
2075 let out_offset = if job.rows() == 0 || job.cols() == 0 {
2077 0
2078 } else {
2079 start - union.start
2080 };
2081 tenferro_tensor::backend::GroupedGemmJob::new(
2082 out_offset,
2083 job.lhs_offset(),
2084 job.rhs_offset(),
2085 job.rows(),
2086 job.contracted(),
2087 job.cols(),
2088 )
2089 })
2090 .collect();
2091 Ok((union, rebased))
2092}
2093
2094#[allow(clippy::too_many_arguments)]
2097fn execute_grouped_outer_typed<T>(
2098 provider: &dyn CpuGemmProvider,
2099 checked: crate::provider::OuterFanOutChecked,
2100 fan_out: crate::provider::CpuOuterFanOut<'_>,
2101 chunks: usize,
2102 lhs: &TensorRead<'_>,
2103 rhs: &TensorRead<'_>,
2104 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
2105 output_storage: &mut [T],
2106 output_base: isize,
2107 wrap_output: for<'a> fn(tenferro_tensor::TypedTensorViewMut<'a, T>) -> TensorViewMut<'a>,
2108) -> Result<()>
2109where
2110 T: Send + Sync + 'static,
2111{
2112 const NO_DUPLICATE: usize = usize::MAX;
2113
2114 let output_base = usize::try_from(output_base).map_err(|_| {
2115 Error::invalid_argument(
2116 "grouped_gemm",
2117 "output",
2118 "grouped-GEMM output base offset is negative",
2119 )
2120 })?;
2121 let output_storage_len = output_storage.len();
2122 let jobs = config.jobs();
2130 let job_count = jobs.len();
2131 let unit_count = if chunks < job_count
2132 && grouped_output_starts_increase(jobs, output_base, output_storage_len)?
2133 {
2134 chunks.max(1)
2135 } else {
2136 for job in jobs {
2137 checked_grouped_output_range(output_base, output_storage_len, job)?;
2138 }
2139 job_count
2140 };
2141 let unit_jobs =
2142 |unit: usize| unit * job_count / unit_count..(unit + 1) * job_count / unit_count;
2143
2144 let output_address = output_storage.as_mut_ptr() as usize;
2145 let operation_error = std::sync::Mutex::new(None);
2146 let failed = std::sync::atomic::AtomicBool::new(false);
2147 let unit_states = PackedJobStates::new(unit_count);
2148 let duplicate_index = AtomicUsize::new(NO_DUPLICATE);
2149 fan_out
2150 .submit(checked, unit_count, |index, provider_context| {
2151 if unit_states.try_claim(index).is_err() {
2152 let _ = duplicate_index.compare_exchange(
2153 NO_DUPLICATE,
2154 index,
2155 Ordering::AcqRel,
2156 Ordering::Acquire,
2157 );
2158 return Err(CpuDomainExecutorError::Scheduling {
2159 message: format!(
2160 "executor invoked grouped-GEMM duplicate index {index}; every index must run exactly once"
2161 ),
2162 });
2163 }
2164
2165 if !failed.load(Ordering::Relaxed) {
2168 let result = (|| -> Result<()> {
2169 let (range, rebased) =
2170 grouped_chunk(&jobs[unit_jobs(index)], output_base, output_storage_len)?;
2171 let start = range.start;
2172 let len = range.len();
2173 let output_slice = unsafe {
2185 std::slice::from_raw_parts_mut((output_address as *mut T).add(start), len)
2186 };
2187 let output_view =
2188 tenferro_tensor::TypedTensorViewMut::from_slice([len], [1], 0, output_slice)?;
2189 let mut output = TensorWrite::from_view(wrap_output(output_view));
2190 let request = CpuGroupedGemmRequest::new(
2191 lhs,
2192 rhs,
2193 &mut output,
2194 &rebased,
2195 config.accumulation(),
2196 );
2197 match provider.grouped_gemm(provider_context, request)? {
2198 CpuProviderOutcome::Executed => Ok(()),
2199 CpuProviderOutcome::Unsupported(reason) => {
2200 Err(unsupported_provider_error("grouped-GEMM", reason))
2201 }
2202 }
2203 })();
2204 if let Err(error) = result {
2205 failed.store(true, Ordering::Relaxed);
2206 *operation_error
2207 .lock()
2208 .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(error);
2209 }
2210 }
2211 let _ = unit_states.complete(index);
2212 Ok(())
2213 })
2214 .map_err(|error| Error::backend_source("grouped_gemm", error))?;
2215 let duplicate = duplicate_index.load(Ordering::Acquire);
2216 if duplicate != NO_DUPLICATE {
2217 return Err(Error::backend_source(
2218 "grouped_gemm",
2219 CpuDomainExecutorError::Scheduling {
2220 message: format!(
2221 "executor invoked grouped-GEMM duplicate index {duplicate}; every index must run exactly once"
2222 ),
2223 },
2224 ));
2225 }
2226 if let Some((index, state)) = unit_states.first_incomplete() {
2227 let detail = if state == GroupedJobState::Unclaimed {
2228 format!("executor omitted grouped-GEMM missing index {index}")
2229 } else {
2230 format!("executor did not complete grouped-GEMM index {index}")
2231 };
2232 return Err(Error::backend_source(
2233 "grouped_gemm",
2234 CpuDomainExecutorError::Scheduling { message: detail },
2235 ));
2236 }
2237 match operation_error.into_inner() {
2238 Ok(Some(error)) => Err(error),
2239 Err(poisoned) => poisoned.into_inner().map_or(Ok(()), Err),
2240 Ok(None) => Ok(()),
2241 }
2242}
2243
2244#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
2253#[error("missing mandatory CPU provider slots: GEMM={gemm}, layout={layout}")]
2254pub struct CpuProviderBundleBuildError {
2255 gemm: bool,
2256 layout: bool,
2257}
2258
2259#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2268pub enum CpuProviderSlot {
2269 Gemm,
2271 LayoutTransform,
2273 GeneralContraction,
2275}
2276
2277#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
2292#[non_exhaustive]
2293pub enum CpuProviderBundleInstallError {
2294 #[error(
2296 "CPU provider bundle slot {provider:?} is incompatible with domain {domain_id:?}: {source}"
2297 )]
2298 IncompatibleDomain {
2299 domain_id: tenferro_tensor::CpuDomainId,
2301 provider: CpuProviderSlot,
2303 #[source]
2305 source: CpuProviderDomainError,
2306 },
2307}
2308
2309#[derive(Debug)]
2320pub struct CpuProviderBundleBuilder {
2321 gemm: Option<Arc<dyn CpuGemmProvider>>,
2322 layout: Option<Arc<dyn CpuLayoutTransformProvider>>,
2323 general: Option<Arc<dyn CpuGeneralContractionProvider>>,
2324 general_policy: GeneralContractionPolicy,
2325 grouped_scheduling: GroupedGemmScheduling,
2326 capability_policy: ProviderCapabilityPolicy,
2327 extensions: crate::provider_extensions::ProviderExtensions,
2328}
2329
2330impl CpuProviderBundleBuilder {
2331 #[cfg(test)]
2332 pub(crate) fn provider_default_compatibility(mut self) -> Self {
2333 self.capability_policy = ProviderCapabilityPolicy::ProviderDefaultCompatibility;
2334 self
2335 }
2336
2337 pub fn gemm_provider(mut self, provider: Arc<dyn CpuGemmProvider>) -> Self {
2339 self.gemm = Some(provider);
2340 self.grouped_scheduling = GroupedGemmScheduling::ProviderOwned;
2341 self
2342 }
2343
2344 pub fn engine_outer_grouped_gemm(mut self) -> Self {
2351 self.grouped_scheduling = GroupedGemmScheduling::EngineOuter;
2352 self
2353 }
2354
2355 pub fn layout_transform_provider(
2357 mut self,
2358 provider: Arc<dyn CpuLayoutTransformProvider>,
2359 ) -> Self {
2360 self.layout = Some(provider);
2361 self
2362 }
2363
2364 pub fn prefer_general_contraction_provider(
2366 mut self,
2367 provider: Arc<dyn CpuGeneralContractionProvider>,
2368 ) -> Self {
2369 self.general = Some(provider);
2370 self.general_policy = GeneralContractionPolicy::Preferred;
2371 self
2372 }
2373
2374 pub fn require_general_contraction_provider(
2376 mut self,
2377 provider: Arc<dyn CpuGeneralContractionProvider>,
2378 ) -> Self {
2379 self.general = Some(provider);
2380 self.general_policy = GeneralContractionPolicy::Required;
2381 self
2382 }
2383
2384 pub fn extension<E: std::any::Any + Send + Sync>(mut self, extension: Arc<E>) -> Self {
2405 self.extensions.insert(extension);
2406 self
2407 }
2408
2409 pub fn build(self) -> std::result::Result<CpuProviderBundle, CpuProviderBundleBuildError> {
2415 let missing = CpuProviderBundleBuildError {
2416 gemm: self.gemm.is_none(),
2417 layout: self.layout.is_none(),
2418 };
2419 let (Some(gemm), Some(layout)) = (self.gemm, self.layout) else {
2420 return Err(missing);
2421 };
2422 let general_capabilities = self
2423 .general
2424 .as_ref()
2425 .map(|provider| provider.execution_capabilities());
2426 let gemm_capabilities = gemm.execution_capabilities();
2427 let layout_capabilities = layout.execution_capabilities();
2428 Ok(CpuProviderBundle {
2429 inner: Arc::new(CpuProviderBundleInner {
2430 dot_general: DotGeneralRuntime {
2431 general: self.general,
2432 gemm,
2433 layout,
2434 general_capabilities,
2435 gemm_capabilities,
2436 layout_capabilities,
2437 general_policy: self.general_policy,
2438 grouped_scheduling: self.grouped_scheduling,
2439 capability_policy: self.capability_policy,
2440 },
2441 extensions: self.extensions,
2442 }),
2443 })
2444 }
2445}
2446
2447fn validate_axis_ranges(axes: &[usize], rank: usize) -> Result<()> {
2448 for &axis in axes {
2449 if axis >= rank {
2450 return Err(Error::axis_out_of_bounds(OP, axis, rank));
2451 }
2452 }
2453 Ok(())
2454}
2455
2456fn role_mask(axes: &[usize], rank: usize, role: &'static str) -> Result<Option<u64>> {
2457 if rank > 64 {
2458 for (position, &axis) in axes.iter().enumerate() {
2459 if axes[..position].contains(&axis) {
2460 return Err(Error::duplicate_axis(OP, axis, role));
2461 }
2462 }
2463 return Ok(None);
2464 }
2465
2466 let mut mask = 0_u64;
2467 for &axis in axes {
2468 let bit = 1_u64 << axis;
2469 if mask & bit != 0 {
2470 return Err(Error::duplicate_axis(OP, axis, role));
2471 }
2472 mask |= bit;
2473 }
2474 Ok(Some(mask))
2475}
2476
2477fn validate_disjoint(
2478 first: &[usize],
2479 first_mask: Option<u64>,
2480 first_role: &'static str,
2481 second: &[usize],
2482 second_mask: Option<u64>,
2483 second_role: &'static str,
2484) -> Result<()> {
2485 let overlap = match (first_mask, second_mask) {
2486 (Some(first), Some(second)) => first & second,
2487 _ => 0,
2488 };
2489 let conflict = if overlap != 0 || first_mask.is_none() {
2490 first.iter().copied().find(|axis| second.contains(axis))
2491 } else {
2492 None
2493 };
2494 if let Some(axis) = conflict {
2495 return Err(Error::validation(
2496 OP,
2497 ValidationError::AxisRoleConflict {
2498 axis,
2499 first_role,
2500 second_role,
2501 },
2502 ));
2503 }
2504 Ok(())
2505}
2506
2507pub(crate) fn validate_axis_groups<'a>(
2508 lhs_rank: usize,
2509 rhs_rank: usize,
2510 config: &'a DotGeneralConfig,
2511) -> Result<CpuContractionAxes<'a>> {
2512 validate_axis_ranges(&config.lhs_contracting_dims, lhs_rank)?;
2513 validate_axis_ranges(&config.rhs_contracting_dims, rhs_rank)?;
2514 validate_axis_ranges(&config.lhs_batch_dims, lhs_rank)?;
2515 validate_axis_ranges(&config.rhs_batch_dims, rhs_rank)?;
2516
2517 let lhs_contracting_mask = role_mask(
2518 &config.lhs_contracting_dims,
2519 lhs_rank,
2520 "lhs_contracting_dims",
2521 )?;
2522 let rhs_contracting_mask = role_mask(
2523 &config.rhs_contracting_dims,
2524 rhs_rank,
2525 "rhs_contracting_dims",
2526 )?;
2527 let lhs_batch_mask = role_mask(&config.lhs_batch_dims, lhs_rank, "lhs_batch_dims")?;
2528 let rhs_batch_mask = role_mask(&config.rhs_batch_dims, rhs_rank, "rhs_batch_dims")?;
2529
2530 validate_disjoint(
2531 &config.lhs_contracting_dims,
2532 lhs_contracting_mask,
2533 "lhs contracting",
2534 &config.lhs_batch_dims,
2535 lhs_batch_mask,
2536 "lhs batch",
2537 )?;
2538 validate_disjoint(
2539 &config.rhs_contracting_dims,
2540 rhs_contracting_mask,
2541 "rhs contracting",
2542 &config.rhs_batch_dims,
2543 rhs_batch_mask,
2544 "rhs batch",
2545 )?;
2546
2547 if config.lhs_contracting_dims.len() != config.rhs_contracting_dims.len() {
2548 return Err(Error::invalid_argument(
2549 OP,
2550 "dot_general_config",
2551 format!(
2552 "lhs/rhs contracting dim counts differ ({} vs {})",
2553 config.lhs_contracting_dims.len(),
2554 config.rhs_contracting_dims.len(),
2555 ),
2556 ));
2557 }
2558 if config.lhs_batch_dims.len() != config.rhs_batch_dims.len() {
2559 return Err(Error::invalid_argument(
2560 OP,
2561 "dot_general_config",
2562 format!(
2563 "lhs/rhs batch dim counts differ ({} vs {})",
2564 config.lhs_batch_dims.len(),
2565 config.rhs_batch_dims.len(),
2566 ),
2567 ));
2568 }
2569
2570 Ok(CpuContractionAxes::new(
2571 lhs_rank,
2572 rhs_rank,
2573 &config.lhs_contracting_dims,
2574 &config.rhs_contracting_dims,
2575 &config.lhs_batch_dims,
2576 &config.rhs_batch_dims,
2577 lhs_contracting_mask.zip(lhs_batch_mask).map(|(a, b)| a | b),
2578 rhs_contracting_mask.zip(rhs_batch_mask).map(|(a, b)| a | b),
2579 ))
2580}
2581
2582#[derive(Clone, Copy, Debug)]
2583pub(crate) struct ValidatedDotGeneral<'a> {
2584 axes: CpuContractionAxes<'a>,
2585 #[cfg(test)]
2586 output_element_count: usize,
2587}
2588
2589impl<'a> ValidatedDotGeneral<'a> {
2590 #[cfg(test)]
2591 pub(crate) fn axes(&self) -> &CpuContractionAxes<'a> {
2592 &self.axes
2593 }
2594
2595 #[cfg(test)]
2596 pub(crate) fn output_element_count(&self) -> usize {
2597 self.output_element_count
2598 }
2599
2600 pub(crate) fn request<'request, 'input, 'output>(
2601 &'request self,
2602 lhs: &'request TensorRead<'input>,
2603 rhs: &'request TensorRead<'input>,
2604 output: &'request mut TensorWrite<'output>,
2605 accumulation: DotGeneralAccumulation,
2606 ) -> CpuDotGeneralRequest<'request, 'input, 'output>
2607 where
2608 'a: 'request,
2609 {
2610 CpuDotGeneralRequest::new(lhs, rhs, output, self.axes, accumulation)
2611 }
2612}
2613
2614fn validate_paired_extents(
2615 lhs: &TensorRead<'_>,
2616 rhs: &TensorRead<'_>,
2617 axes: &CpuContractionAxes<'_>,
2618) -> Result<()> {
2619 for (lhs_axis, rhs_axis) in axes.contracting_pairs().chain(axes.batch_pairs()) {
2620 if lhs.shape()[lhs_axis] != rhs.shape()[rhs_axis] {
2621 return Err(Error::validation(
2622 OP,
2623 ShapeMismatch::ContractedDimensions {
2624 lhs_axis,
2625 lhs_size: lhs.shape()[lhs_axis],
2626 rhs_axis,
2627 rhs_size: rhs.shape()[rhs_axis],
2628 }
2629 .into(),
2630 ));
2631 }
2632 }
2633 Ok(())
2634}
2635
2636fn expected_output_shape(
2637 lhs: &TensorRead<'_>,
2638 rhs: &TensorRead<'_>,
2639 axes: &CpuContractionAxes<'_>,
2640) -> Vec<usize> {
2641 axes.lhs_free_axes()
2642 .map(|axis| lhs.shape()[axis])
2643 .chain(axes.rhs_free_axes().map(|axis| rhs.shape()[axis]))
2644 .chain(
2645 axes.batch_pairs()
2646 .map(|(lhs_axis, _)| lhs.shape()[lhs_axis]),
2647 )
2648 .collect()
2649}
2650
2651fn output_shape_matches(
2652 lhs: &TensorRead<'_>,
2653 rhs: &TensorRead<'_>,
2654 output: &TensorWrite<'_>,
2655 axes: &CpuContractionAxes<'_>,
2656) -> Result<()> {
2657 let expected_rank =
2658 axes.lhs_free_axes().count() + axes.rhs_free_axes().count() + axes.batch_pairs().len();
2659 let mut actual = output.shape().iter().copied();
2660 let matches = output.shape().len() == expected_rank
2661 && axes
2662 .lhs_free_axes()
2663 .map(|axis| lhs.shape()[axis])
2664 .chain(axes.rhs_free_axes().map(|axis| rhs.shape()[axis]))
2665 .chain(
2666 axes.batch_pairs()
2667 .map(|(lhs_axis, _)| lhs.shape()[lhs_axis]),
2668 )
2669 .all(|expected| actual.next() == Some(expected));
2670 if matches {
2671 return Ok(());
2672 }
2673
2674 Err(Error::validation(
2675 OP,
2676 ShapeMismatch::ExpectedActual {
2677 expected: expected_output_shape(lhs, rhs, axes).into(),
2678 actual: output.shape().to_vec().into(),
2679 }
2680 .into(),
2681 ))
2682}
2683
2684fn layout_overflow() -> Error {
2685 Error::validation(OP, ValidationError::IntegerOverflow)
2686}
2687
2688pub(crate) fn validate_layout_metadata(
2689 role: &'static str,
2690 shape: &[usize],
2691 strides: &[isize],
2692 offset: isize,
2693 storage_len: usize,
2694) -> Result<usize> {
2695 if shape.len() != strides.len() {
2696 return Err(Error::validation(
2697 OP,
2698 ValidationError::RankMismatch {
2699 expected: shape.len(),
2700 actual: strides.len(),
2701 },
2702 ));
2703 }
2704 let element_count = tenferro_tensor::validate::checked_shape_product(OP, role, shape)?;
2705
2706 if shape.contains(&0) {
2707 let offset = usize::try_from(offset).map_err(|_| {
2708 Error::invalid_argument(OP, role, "minimum reachable offset is negative")
2709 })?;
2710 if offset > storage_len {
2711 return Err(Error::validation(OP, ValidationError::ViewOutOfBounds));
2712 }
2713 return Ok(element_count);
2714 }
2715
2716 let mut minimum = offset;
2717 let mut maximum = offset;
2718 for (&extent, &stride) in shape.iter().zip(strides) {
2719 let steps = isize::try_from(extent - 1).map_err(|_| layout_overflow())?;
2720 let end = stride.checked_mul(steps).ok_or_else(layout_overflow)?;
2721 let (axis_minimum, axis_maximum) = if end < 0 { (end, 0) } else { (0, end) };
2722 minimum = minimum
2723 .checked_add(axis_minimum)
2724 .ok_or_else(layout_overflow)?;
2725 maximum = maximum
2726 .checked_add(axis_maximum)
2727 .ok_or_else(layout_overflow)?;
2728 }
2729 let minimum = usize::try_from(minimum)
2730 .map_err(|_| Error::invalid_argument(OP, role, "minimum reachable offset is negative"))?;
2731 let maximum = usize::try_from(maximum)
2732 .map_err(|_| Error::invalid_argument(OP, role, "maximum reachable offset is negative"))?;
2733 if minimum > maximum || maximum >= storage_len {
2734 return Err(Error::validation(OP, ValidationError::ViewOutOfBounds));
2735 }
2736 Ok(element_count)
2737}
2738
2739macro_rules! validate_owned_layout {
2740 ($tensor:expr, $role:expr) => {{
2741 let tensor = $tensor;
2742 if tensor.backend_buffer().is_some() {
2743 return Err(crate::cpu_backend_buffer_error(OP));
2744 }
2745 let storage_len = tensor.host_data()?.len();
2749 let element_count =
2750 tenferro_tensor::validate::checked_shape_product(OP, $role, tensor.shape())?;
2751 if element_count > storage_len {
2752 return Err(Error::validation(OP, ValidationError::ViewOutOfBounds));
2753 }
2754 Ok(element_count)
2755 }};
2756}
2757
2758macro_rules! validate_read_view_layout {
2759 ($view:expr, $role:expr) => {{
2760 let view = $view;
2761 let storage_len = view.host_storage()?.len();
2762 validate_layout_metadata(
2763 $role,
2764 view.shape(),
2765 view.strides(),
2766 view.offset(),
2767 storage_len,
2768 )
2769 }};
2770}
2771
2772macro_rules! validate_write_view_layout {
2773 ($view:expr, $role:expr) => {{
2774 let view = $view;
2775 let storage_len = view.host_storage()?.len();
2776 validate_layout_metadata(
2777 $role,
2778 view.shape(),
2779 view.strides(),
2780 view.offset(),
2781 storage_len,
2782 )
2783 }};
2784}
2785
2786macro_rules! validate_owned_layout_table {
2793 ($tensor:expr, $op:expr, $message:expr, $role:expr) => {
2794 match $tensor.dtype() {
2795 DType::External(type_id) => Err(crate::Error::unsupported_dtype(
2797 $op,
2798 DType::External(type_id),
2799 $message,
2800 )),
2801 DType::F32 => {
2802 validate_owned_layout!(validated_operand::<f32>($tensor, $op, $message)?, $role)
2803 }
2804 DType::F64 => {
2805 validate_owned_layout!(validated_operand::<f64>($tensor, $op, $message)?, $role)
2806 }
2807 DType::I32 => {
2808 validate_owned_layout!(validated_operand::<i32>($tensor, $op, $message)?, $role)
2809 }
2810 DType::I64 => {
2811 validate_owned_layout!(validated_operand::<i64>($tensor, $op, $message)?, $role)
2812 }
2813 DType::Bool => {
2814 validate_owned_layout!(validated_operand::<bool>($tensor, $op, $message)?, $role)
2815 }
2816 DType::C32 => validate_owned_layout!(
2817 validated_operand::<Complex32>($tensor, $op, $message)?,
2818 $role
2819 ),
2820 DType::C64 => validate_owned_layout!(
2821 validated_operand::<Complex64>($tensor, $op, $message)?,
2822 $role
2823 ),
2824 }
2825 };
2826}
2827
2828fn validate_read_layout(tensor: &TensorRead<'_>, role: &'static str) -> Result<usize> {
2829 match tensor {
2830 TensorRead::Tensor(tensor) => validate_owned_layout_table!(
2831 tensor,
2832 "validate_read_layout",
2833 "an externally defined payload is not a runtime operand",
2834 role
2835 ),
2836 TensorRead::View(view) => match view {
2837 TensorView::F32(view) => validate_read_view_layout!(view, role),
2838 TensorView::F64(view) => validate_read_view_layout!(view, role),
2839 TensorView::I32(view) => validate_read_view_layout!(view, role),
2840 TensorView::I64(view) => validate_read_view_layout!(view, role),
2841 TensorView::Bool(view) => validate_read_view_layout!(view, role),
2842 TensorView::C32(view) => validate_read_view_layout!(view, role),
2843 TensorView::C64(view) => validate_read_view_layout!(view, role),
2844 },
2845 }
2846}
2847
2848fn validate_write_layout(tensor: &TensorWrite<'_>, role: &'static str) -> Result<usize> {
2849 match tensor {
2850 TensorWrite::Tensor(tensor) => validate_owned_layout_table!(
2851 tensor,
2852 "validate_write_layout",
2853 "an externally defined payload is not a runtime destination",
2854 role
2855 ),
2856 TensorWrite::View(view) => match view {
2857 TensorViewMut::F32(view) => validate_write_view_layout!(view, role),
2858 TensorViewMut::F64(view) => validate_write_view_layout!(view, role),
2859 TensorViewMut::I32(view) => validate_write_view_layout!(view, role),
2860 TensorViewMut::I64(view) => validate_write_view_layout!(view, role),
2861 TensorViewMut::Bool(view) => validate_write_view_layout!(view, role),
2862 TensorViewMut::C32(view) => validate_write_view_layout!(view, role),
2863 TensorViewMut::C64(view) => validate_write_view_layout!(view, role),
2864 },
2865 }
2866}
2867
2868pub(crate) fn validate_dot_general<'a>(
2869 lhs: &TensorRead<'_>,
2870 rhs: &TensorRead<'_>,
2871 output: &TensorWrite<'_>,
2872 config: &'a DotGeneralConfig,
2873 accumulation: DotGeneralAccumulation,
2874) -> Result<ValidatedDotGeneral<'a>> {
2875 if lhs.dtype() != rhs.dtype() {
2876 return Err(Error::dtype_mismatch(OP, lhs.dtype(), rhs.dtype()));
2877 }
2878 if output.dtype() != lhs.dtype() {
2879 return Err(Error::dtype_mismatch(OP, output.dtype(), lhs.dtype()));
2880 }
2881 if accumulation.alpha.dtype() != lhs.dtype() {
2882 return Err(Error::dtype_mismatch(
2883 OP,
2884 lhs.dtype(),
2885 accumulation.alpha.dtype(),
2886 ));
2887 }
2888 if accumulation.beta.dtype() != lhs.dtype() {
2889 return Err(Error::dtype_mismatch(
2890 OP,
2891 lhs.dtype(),
2892 accumulation.beta.dtype(),
2893 ));
2894 }
2895
2896 crate::structural::validate_cpu_host_placement(OP, "lhs", read_placement(lhs))?;
2897 crate::structural::validate_cpu_host_placement(OP, "rhs", read_placement(rhs))?;
2898 crate::structural::validate_cpu_host_placement(OP, "output", write_placement(output))?;
2899 validate_read_layout(lhs, "lhs")?;
2900 validate_read_layout(rhs, "rhs")?;
2901 let _output_element_count = validate_write_layout(output, "output")?;
2903
2904 let axes = validate_axis_groups(lhs.shape().len(), rhs.shape().len(), config)?;
2905 validate_paired_extents(lhs, rhs, &axes)?;
2906 output_shape_matches(lhs, rhs, output, &axes)?;
2907
2908 Ok(ValidatedDotGeneral {
2909 axes,
2910 #[cfg(test)]
2911 output_element_count: _output_element_count,
2912 })
2913}
2914
2915fn read_placement<'a>(tensor: &'a TensorRead<'_>) -> &'a tenferro_tensor::Placement {
2916 match tensor {
2917 TensorRead::Tensor(tensor) => tensor.placement(),
2918 TensorRead::View(view) => match view {
2919 tenferro_tensor::TensorView::F32(view) => view.placement(),
2920 tenferro_tensor::TensorView::F64(view) => view.placement(),
2921 tenferro_tensor::TensorView::I32(view) => view.placement(),
2922 tenferro_tensor::TensorView::I64(view) => view.placement(),
2923 tenferro_tensor::TensorView::Bool(view) => view.placement(),
2924 tenferro_tensor::TensorView::C32(view) => view.placement(),
2925 tenferro_tensor::TensorView::C64(view) => view.placement(),
2926 },
2927 }
2928}
2929
2930fn write_placement<'a>(tensor: &'a TensorWrite<'_>) -> &'a tenferro_tensor::Placement {
2931 match tensor {
2932 TensorWrite::Tensor(tensor) => tensor.placement(),
2933 TensorWrite::View(view) => match view {
2934 tenferro_tensor::TensorViewMut::F32(view) => view.placement(),
2935 tenferro_tensor::TensorViewMut::F64(view) => view.placement(),
2936 tenferro_tensor::TensorViewMut::I32(view) => view.placement(),
2937 tenferro_tensor::TensorViewMut::I64(view) => view.placement(),
2938 tenferro_tensor::TensorViewMut::Bool(view) => view.placement(),
2939 tenferro_tensor::TensorViewMut::C32(view) => view.placement(),
2940 tenferro_tensor::TensorViewMut::C64(view) => view.placement(),
2941 },
2942 }
2943}
2944
2945#[cfg(test)]
2946mod tests;