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::buffer_pool::{BufferPool, PoolScalar};
14use crate::provider::{
15 builtin_gemm_provider, builtin_layout_provider, CpuContractionAxes, CpuDotGeneralRequest,
16 CpuExecutionContext, CpuGemmProvider, CpuGeneralContractionProvider, CpuGroupedGemmRequest,
17 CpuLayoutTransformIntent, CpuLayoutTransformProvider, CpuLayoutTransformRequest,
18 CpuOperationEntry, CpuProviderOutcome, CpuProviderUnsupported, CpuUninitGemmProvider,
19};
20use crate::{
21 gemm::GemmAnalysisCache, CpuDomainExecutorError, CpuDomainId, CpuPlacementGuarantee,
22 CpuProviderDomainError, CpuSet, Error, ParallelMode, PooledUninitOutput, Result,
23};
24
25const OP: &str = "dot_general";
26
27#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
40pub enum GeneralContractionPolicy {
41 #[default]
43 Preferred,
44 Required,
46}
47
48#[derive(Debug)]
49pub(crate) struct DotGeneralRuntime {
50 pub(crate) general: Option<Arc<dyn CpuGeneralContractionProvider>>,
51 pub(crate) gemm: Arc<dyn CpuGemmProvider>,
52 pub(crate) layout: Arc<dyn CpuLayoutTransformProvider>,
53 general_capabilities: Option<crate::CpuProviderExecutionCapabilities>,
54 gemm_capabilities: crate::CpuProviderExecutionCapabilities,
55 layout_capabilities: crate::CpuProviderExecutionCapabilities,
56 pub(crate) general_policy: GeneralContractionPolicy,
57 grouped_scheduling: GroupedGemmScheduling,
58 capability_policy: ProviderCapabilityPolicy,
59}
60
61#[derive(Clone, Copy, Debug, PartialEq, Eq)]
62enum GroupedGemmScheduling {
63 ProviderOwned,
64 EngineOuter,
65}
66
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68enum ProviderCapabilityPolicy {
69 Strict,
70 ProviderDefaultCompatibility,
71}
72
73const GROUPED_JOB_STATE_BITS: usize = 2;
74const GROUPED_JOBS_PER_STATE_WORD: usize = usize::BITS as usize / GROUPED_JOB_STATE_BITS;
75const GROUPED_INLINE_STATE_WORDS: usize = 4;
76const GROUPED_INLINE_JOB_CAPACITY: usize = GROUPED_INLINE_STATE_WORDS * GROUPED_JOBS_PER_STATE_WORD;
77
78#[derive(Clone, Copy, Debug, Eq, PartialEq)]
79#[repr(usize)]
80enum GroupedJobState {
81 Unclaimed = 0,
82 Running = 1,
83 Complete = 2,
84 Reserved = 3,
85}
86
87impl GroupedJobState {
88 fn from_bits(bits: usize) -> Self {
89 match bits {
90 0 => Self::Unclaimed,
91 1 => Self::Running,
92 2 => Self::Complete,
93 _ => Self::Reserved,
94 }
95 }
96}
97
98struct PackedJobStates {
104 words: SmallVec<[AtomicUsize; GROUPED_INLINE_STATE_WORDS]>,
105 len: usize,
106}
107
108impl PackedJobStates {
109 fn new(len: usize) -> Self {
110 let word_count = len.div_ceil(GROUPED_JOBS_PER_STATE_WORD);
111 let mut words = SmallVec::new();
112 words.resize_with(word_count, || AtomicUsize::new(0));
113 Self { words, len }
114 }
115
116 fn position(index: usize) -> (usize, usize) {
117 let word = index / GROUPED_JOBS_PER_STATE_WORD;
118 let shift = (index % GROUPED_JOBS_PER_STATE_WORD) * GROUPED_JOB_STATE_BITS;
119 (word, shift)
120 }
121
122 fn state(&self, index: usize) -> GroupedJobState {
123 let (word, shift) = Self::position(index);
124 let bits = (self.words[word].load(Ordering::Acquire) >> shift) & 0b11;
125 GroupedJobState::from_bits(bits)
126 }
127
128 fn try_claim(&self, index: usize) -> std::result::Result<(), GroupedJobState> {
129 let (word, shift) = Self::position(index);
130 let word = &self.words[word];
131 let mask = 0b11usize << shift;
132 let running = (GroupedJobState::Running as usize) << shift;
133 let mut observed = word.load(Ordering::Acquire);
134 loop {
135 let state = GroupedJobState::from_bits((observed & mask) >> shift);
136 if state != GroupedJobState::Unclaimed {
137 return Err(state);
138 }
139 let updated = (observed & !mask) | running;
140 match word.compare_exchange_weak(observed, updated, Ordering::AcqRel, Ordering::Acquire)
141 {
142 Ok(_) => return Ok(()),
143 Err(current) => observed = current,
144 }
145 }
146 }
147
148 fn complete(&self, index: usize) -> bool {
149 let (word, shift) = Self::position(index);
150 let word = &self.words[word];
151 let mask = 0b11usize << shift;
152 let complete = (GroupedJobState::Complete as usize) << shift;
153 let mut observed = word.load(Ordering::Acquire);
154 loop {
155 if GroupedJobState::from_bits((observed & mask) >> shift) != GroupedJobState::Running {
156 return false;
157 }
158 let updated = (observed & !mask) | complete;
159 match word.compare_exchange_weak(observed, updated, Ordering::AcqRel, Ordering::Acquire)
160 {
161 Ok(_) => return true,
162 Err(current) => observed = current,
163 }
164 }
165 }
166
167 fn first_incomplete(&self) -> Option<(usize, GroupedJobState)> {
168 (0..self.len)
169 .map(|index| (index, self.state(index)))
170 .find(|(_, state)| *state != GroupedJobState::Complete)
171 }
172
173 #[cfg(test)]
174 fn len(&self) -> usize {
175 self.len
176 }
177
178 #[cfg(test)]
179 fn word_count(&self) -> usize {
180 self.words.len()
181 }
182
183 #[cfg(test)]
184 fn spilled(&self) -> bool {
185 self.words.spilled()
186 }
187}
188
189fn standard_grouped_scheduling(kind: CpuBackendKind) -> GroupedGemmScheduling {
190 match kind {
191 CpuBackendKind::Faer => GroupedGemmScheduling::EngineOuter,
192 CpuBackendKind::Blas => GroupedGemmScheduling::ProviderOwned,
193 }
194}
195
196#[derive(Debug)]
197pub(crate) struct CpuProviderBundleInner {
198 pub(crate) dot_general: DotGeneralRuntime,
199}
200
201#[derive(Clone, Copy)]
202pub(crate) enum CpuProviderDomainContract<'a> {
203 CooperativeCpuSet {
204 placement_guarantee: CpuPlacementGuarantee,
205 domain_cpus: &'a CpuSet,
206 process_allowed_cpus: &'a CpuSet,
207 },
208 CallerManaged,
209}
210
211#[derive(Clone, Debug)]
226pub struct CpuProviderBundle {
227 inner: Arc<CpuProviderBundleInner>,
228}
229
230impl CpuProviderBundle {
231 pub(crate) fn standard(kind: CpuBackendKind, provider_default_compatibility: bool) -> Self {
232 let gemm = builtin_gemm_provider(kind);
233 let layout = builtin_layout_provider();
234 let gemm_capabilities = gemm.execution_capabilities();
235 let layout_capabilities = layout.execution_capabilities();
236 Self {
237 inner: Arc::new(CpuProviderBundleInner {
238 dot_general: DotGeneralRuntime {
239 general: None,
240 gemm,
241 layout,
242 general_capabilities: None,
243 gemm_capabilities,
244 layout_capabilities,
245 general_policy: GeneralContractionPolicy::Preferred,
246 grouped_scheduling: standard_grouped_scheduling(kind),
247 capability_policy: if provider_default_compatibility {
248 ProviderCapabilityPolicy::ProviderDefaultCompatibility
249 } else {
250 ProviderCapabilityPolicy::Strict
251 },
252 },
253 }),
254 }
255 }
256
257 pub fn builder(kind: CpuBackendKind) -> CpuProviderBundleBuilder {
259 CpuProviderBundleBuilder {
260 gemm: Some(builtin_gemm_provider(kind)),
261 layout: Some(builtin_layout_provider()),
262 general: None,
263 general_policy: GeneralContractionPolicy::Preferred,
264 grouped_scheduling: standard_grouped_scheduling(kind),
265 capability_policy: ProviderCapabilityPolicy::Strict,
266 }
267 }
268
269 pub fn custom_builder() -> CpuProviderBundleBuilder {
271 CpuProviderBundleBuilder {
272 gemm: None,
273 layout: None,
274 general: None,
275 general_policy: GeneralContractionPolicy::Preferred,
276 grouped_scheduling: GroupedGemmScheduling::ProviderOwned,
277 capability_policy: ProviderCapabilityPolicy::Strict,
278 }
279 }
280
281 pub fn shares_identity_with(&self, other: &Self) -> bool {
283 Arc::ptr_eq(&self.inner, &other.inner)
284 }
285
286 pub(crate) fn inner(&self) -> &Arc<CpuProviderBundleInner> {
287 &self.inner
288 }
289
290 pub(crate) fn dot_general(&self) -> &DotGeneralRuntime {
291 &self.inner.dot_general
292 }
293
294 pub(crate) fn validate_for_domain(
295 &self,
296 domain_id: CpuDomainId,
297 thread_budget: std::num::NonZeroUsize,
298 contract: CpuProviderDomainContract<'_>,
299 ) -> std::result::Result<(), CpuProviderBundleInstallError> {
300 let runtime = self.dot_general();
301 let validate = |provider, capabilities| {
302 let result = match contract {
303 CpuProviderDomainContract::CooperativeCpuSet {
304 placement_guarantee,
305 domain_cpus,
306 process_allowed_cpus,
307 } => crate::provider_capability::validate_provider_for_domain(
308 capabilities,
309 thread_budget,
310 placement_guarantee,
311 domain_cpus,
312 process_allowed_cpus,
313 ),
314 CpuProviderDomainContract::CallerManaged => {
315 crate::provider_capability::validate_provider_for_caller_managed_domain(
316 capabilities,
317 thread_budget,
318 )
319 }
320 };
321 result.map_err(|source| CpuProviderBundleInstallError::IncompatibleDomain {
322 domain_id,
323 provider,
324 source,
325 })
326 };
327
328 if let Some(capabilities) = runtime.general_capabilities {
329 validate(CpuProviderSlot::GeneralContraction, capabilities)?;
330 }
331 validate(CpuProviderSlot::Gemm, runtime.gemm_capabilities)?;
332 validate(
333 CpuProviderSlot::LayoutTransform,
334 runtime.layout_capabilities,
335 )?;
336
337 let selected_mode = if thread_budget.get() == 1 {
338 ParallelMode::Sequential
339 } else if runtime.accepts_dot_general_mode(ParallelMode::Inner) {
340 ParallelMode::Inner
341 } else {
342 ParallelMode::Sequential
343 };
344 for (provider, capabilities) in [
345 (CpuProviderSlot::Gemm, runtime.gemm_capabilities),
346 (
347 CpuProviderSlot::LayoutTransform,
348 runtime.layout_capabilities,
349 ),
350 ] {
351 if !capabilities.accepts_mode(selected_mode) {
352 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
353 domain_id,
354 provider,
355 source: CpuProviderDomainError::ParallelModeNotSupported {
356 mode: selected_mode,
357 },
358 });
359 }
360 }
361 if let Some(capabilities) = runtime.general_capabilities {
362 if !capabilities.accepts_mode(selected_mode) {
363 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
364 domain_id,
365 provider: CpuProviderSlot::GeneralContraction,
366 source: CpuProviderDomainError::ParallelModeNotSupported {
367 mode: selected_mode,
368 },
369 });
370 }
371 }
372 if runtime.grouped_scheduling == GroupedGemmScheduling::EngineOuter
373 && !runtime.gemm_capabilities.accepts_mode(ParallelMode::Outer)
374 {
375 return Err(CpuProviderBundleInstallError::IncompatibleDomain {
376 domain_id,
377 provider: CpuProviderSlot::Gemm,
378 source: CpuProviderDomainError::ParallelModeNotSupported {
379 mode: ParallelMode::Outer,
380 },
381 });
382 }
383 Ok(())
384 }
385
386 pub(crate) fn preflight_dot_general(&self, entry: &CpuOperationEntry<'_>) -> Result<()> {
387 self.inner
388 .dot_general
389 .dot_general_mode(entry)
390 .map(|_| ())
391 .map_err(|error| Error::backend_source(OP, error))
392 }
393
394 #[allow(clippy::too_many_arguments)]
395 pub(crate) fn execute_dot_general_into(
396 &self,
397 entry: &CpuOperationEntry<'_>,
398 buffers: &mut BufferPool,
399 cache: &mut GemmAnalysisCache,
400 cache_slot: Option<usize>,
401 lhs: TensorRead<'_>,
402 rhs: TensorRead<'_>,
403 config: &DotGeneralConfig,
404 accumulation: DotGeneralAccumulation,
405 output: TensorWrite<'_>,
406 ) -> Result<()> {
407 self.execute_dot_general_into_scoped(
408 entry,
409 None,
410 buffers,
411 cache,
412 cache_slot,
413 lhs,
414 rhs,
415 config,
416 accumulation,
417 output,
418 )
419 }
420
421 #[allow(clippy::too_many_arguments)]
422 pub(crate) fn execute_dot_general_into_scoped(
423 &self,
424 entry: &CpuOperationEntry<'_>,
425 entered: Option<&CpuExecutionContext<'_>>,
426 buffers: &mut BufferPool,
427 cache: &mut GemmAnalysisCache,
428 cache_slot: Option<usize>,
429 lhs: TensorRead<'_>,
430 rhs: TensorRead<'_>,
431 config: &DotGeneralConfig,
432 accumulation: DotGeneralAccumulation,
433 output: TensorWrite<'_>,
434 ) -> Result<()> {
435 self.inner.dot_general.execute_into(
436 &self.inner,
437 entry,
438 entered,
439 buffers,
440 cache,
441 cache_slot,
442 lhs,
443 rhs,
444 config,
445 accumulation,
446 output,
447 )
448 }
449
450 pub(crate) fn execute_grouped_gemm(
451 &self,
452 entry: &CpuOperationEntry<'_>,
453 lhs: TensorRead<'_>,
454 rhs: TensorRead<'_>,
455 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
456 output: TensorWrite<'_>,
457 ) -> Result<()> {
458 self.execute_grouped_gemm_scoped(entry, None, lhs, rhs, config, output)
459 }
460
461 pub(crate) fn execute_grouped_gemm_scoped(
462 &self,
463 entry: &CpuOperationEntry<'_>,
464 entered: Option<&CpuExecutionContext<'_>>,
465 lhs: TensorRead<'_>,
466 rhs: TensorRead<'_>,
467 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
468 output: TensorWrite<'_>,
469 ) -> Result<()> {
470 self.inner
471 .dot_general
472 .execute_grouped(entry, entered, lhs, rhs, config, output)
473 }
474}
475
476fn unsupported_provider_error(capability: &'static str, reason: CpuProviderUnsupported) -> Error {
477 Error::unsupported(
478 OP,
479 format!("configured CPU {capability} provider reported unsupported: {reason:?}"),
480 )
481}
482
483impl DotGeneralRuntime {
484 fn accepts_dot_general_mode(&self, mode: crate::ParallelMode) -> bool {
485 self.general_capabilities
486 .is_none_or(|capabilities| capabilities.accepts_mode(mode))
487 && self.gemm_capabilities.accepts_mode(mode)
488 && self.layout_capabilities.accepts_mode(mode)
489 }
490
491 fn validate_strict_capability(
492 &self,
493 capabilities: crate::CpuProviderExecutionCapabilities,
494 thread_budget: usize,
495 ) -> std::result::Result<(), CpuProviderDomainError> {
496 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
497 return Ok(());
498 }
499 if capabilities.thread_count == crate::CpuThreadCountControl::GlobalOrUncontrolled {
500 return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
501 thread_budget,
502 control: capabilities.thread_count,
503 });
504 }
505 Ok(())
506 }
507
508 fn dot_general_mode(
509 &self,
510 entry: &CpuOperationEntry<'_>,
511 ) -> std::result::Result<ParallelMode, CpuProviderDomainError> {
512 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
513 return Ok(entry.provider_default_compatibility_mode());
514 }
515 let thread_budget = entry.thread_budget().get();
516 if let Some(capabilities) = self.general_capabilities {
517 self.validate_strict_capability(capabilities, thread_budget)?;
518 }
519 self.validate_strict_capability(self.gemm_capabilities, thread_budget)?;
520 self.validate_strict_capability(self.layout_capabilities, thread_budget)?;
521 entry.preferred_provider_mode(|mode| self.accepts_dot_general_mode(mode))
522 }
523
524 fn grouped_mode(
525 &self,
526 entry: &CpuOperationEntry<'_>,
527 ) -> std::result::Result<ParallelMode, CpuProviderDomainError> {
528 if self.capability_policy == ProviderCapabilityPolicy::ProviderDefaultCompatibility {
529 return Ok(entry.provider_default_compatibility_mode());
530 }
531 self.validate_strict_capability(self.gemm_capabilities, entry.thread_budget().get())?;
532 entry.preferred_provider_mode(|mode| self.gemm_capabilities.accepts_mode(mode))
533 }
534
535 #[allow(clippy::too_many_arguments)]
536 fn execute_into(
537 &self,
538 bundle_identity: &Arc<CpuProviderBundleInner>,
539 entry: &CpuOperationEntry<'_>,
540 entered: Option<&CpuExecutionContext<'_>>,
541 buffers: &mut BufferPool,
542 cache: &mut GemmAnalysisCache,
543 cache_slot: Option<usize>,
544 lhs: TensorRead<'_>,
545 rhs: TensorRead<'_>,
546 config: &DotGeneralConfig,
547 accumulation: DotGeneralAccumulation,
548 output: TensorWrite<'_>,
549 ) -> Result<()> {
550 let validated = validate_dot_general(&lhs, &rhs, &output, config, accumulation)?;
551 let mode = self
552 .dot_general_mode(entry)
553 .map_err(|error| Error::backend_source(OP, error))?;
554 cache.bind_provider_bundle(bundle_identity);
555 entry
556 .enter_or_reuse(entered, mode, |provider_context| {
557 self.execute_into_validated(
558 provider_context,
559 validated,
560 buffers,
561 cache,
562 cache_slot,
563 lhs,
564 rhs,
565 config,
566 accumulation,
567 output,
568 )
569 })
570 .map_err(|error| Error::backend_source(OP, error))?
571 }
572
573 #[allow(clippy::too_many_arguments)]
577 fn execute_into_validated(
578 &self,
579 provider_context: &CpuExecutionContext<'_>,
580 validated: ValidatedDotGeneral<'_>,
581 buffers: &mut BufferPool,
582 cache: &mut GemmAnalysisCache,
583 cache_slot: Option<usize>,
584 lhs: TensorRead<'_>,
585 rhs: TensorRead<'_>,
586 config: &DotGeneralConfig,
587 accumulation: DotGeneralAccumulation,
588 mut output: TensorWrite<'_>,
589 ) -> Result<()> {
590 if let Some(general) = &self.general {
591 let request = validated.request(&lhs, &rhs, &mut output, accumulation);
592 match general.dot_general(provider_context, request)? {
593 CpuProviderOutcome::Executed => return Ok(()),
594 CpuProviderOutcome::Unsupported(reason) => {
595 if self.general_policy == GeneralContractionPolicy::Required {
596 return Err(unsupported_provider_error(
597 "required general-contraction",
598 reason,
599 ));
600 }
601 }
602 }
603 }
604
605 if let Some(plan) =
606 crate::gemm::prepare_provider_gemm(cache, cache_slot, &lhs, &rhs, &output, config)?
607 {
608 match execute_gemm_plan(
609 self.gemm.as_ref(),
610 provider_context,
611 plan,
612 &lhs,
613 &rhs,
614 accumulation,
615 &mut output,
616 )? {
617 CpuProviderOutcome::Executed => return Ok(()),
618 CpuProviderOutcome::Unsupported(reason)
619 if !canonical_gemm_fallback_supported(reason) =>
620 {
621 return Err(unsupported_provider_error("GEMM", reason));
622 }
623 CpuProviderOutcome::Unsupported(_) => {}
624 }
625 }
626
627 self.execute_canonical_gemm(
628 provider_context,
629 buffers,
630 cache,
631 cache_slot,
632 &lhs,
633 &rhs,
634 config,
635 accumulation,
636 &mut output,
637 )
638 }
639
640 #[allow(clippy::too_many_arguments)]
641 fn execute_canonical_gemm(
642 &self,
643 provider_context: &CpuExecutionContext<'_>,
644 buffers: &mut BufferPool,
645 cache: &mut GemmAnalysisCache,
646 cache_slot: Option<usize>,
647 lhs: &TensorRead<'_>,
648 rhs: &TensorRead<'_>,
649 config: &DotGeneralConfig,
650 accumulation: DotGeneralAccumulation,
651 output: &mut TensorWrite<'_>,
652 ) -> Result<()> {
653 let (lhs_perm, rhs_perm, canonical_config) =
654 crate::gemm::canonical_gemm_layout(config, lhs.shape().len(), rhs.shape().len());
655 let lhs_canonical = materialize_canonical_operand(
656 self.layout.as_ref(),
657 provider_context,
658 buffers,
659 lhs,
660 &lhs_perm,
661 accumulation.lhs_conj,
662 )?;
663 let rhs_canonical = match materialize_canonical_operand(
664 self.layout.as_ref(),
665 provider_context,
666 buffers,
667 rhs,
668 &rhs_perm,
669 accumulation.rhs_conj,
670 ) {
671 Ok(tensor) => tensor,
672 Err(error) => {
673 reclaim_temporary(buffers, lhs_canonical);
674 return Err(error);
675 }
676 };
677
678 let result = {
679 let lhs = TensorRead::from_tensor(&lhs_canonical);
680 let rhs = TensorRead::from_tensor(&rhs_canonical);
681 let canonical_accumulation = DotGeneralAccumulation {
682 lhs_conj: false,
683 rhs_conj: false,
684 ..accumulation
685 };
686 match crate::gemm::prepare_provider_gemm_canonical(
687 cache,
688 cache_slot,
689 &lhs,
690 &rhs,
691 output,
692 &canonical_config,
693 ) {
694 Ok(Some(plan)) => match execute_gemm_plan(
695 self.gemm.as_ref(),
696 provider_context,
697 plan,
698 &lhs,
699 &rhs,
700 canonical_accumulation,
701 output,
702 ) {
703 Ok(CpuProviderOutcome::Executed) => Ok(()),
704 Ok(CpuProviderOutcome::Unsupported(reason)) => {
705 Err(unsupported_provider_error("GEMM", reason))
706 }
707 Err(error) => Err(error),
708 },
709 Ok(None) => Err(Error::unsupported(
710 OP,
711 "configured CPU layout-plus-GEMM path cannot represent the canonical contraction",
712 )),
713 Err(error) => Err(error),
714 }
715 };
716 reclaim_temporary(buffers, lhs_canonical);
717 reclaim_temporary(buffers, rhs_canonical);
718 result
719 }
720
721 #[allow(clippy::too_many_arguments)]
734 pub(crate) fn execute_dot_into_uninit(
735 &self,
736 bundle_identity: &Arc<CpuProviderBundleInner>,
737 entry: &CpuOperationEntry<'_>,
738 entered: Option<&CpuExecutionContext<'_>>,
739 cache: &mut GemmAnalysisCache,
740 cache_slot: Option<usize>,
741 lhs: &TensorRead<'_>,
742 rhs: &TensorRead<'_>,
743 config: &DotGeneralConfig,
744 accumulation: DotGeneralAccumulation,
745 output_shape: &[usize],
746 output_bytes: &mut [MaybeUninit<u8>],
747 ) -> Result<CpuProviderOutcome> {
748 let Some(witness) = self.gemm.uninit_provider() else {
749 return Err(Error::unsupported(
750 OP,
751 "configured CPU GEMM provider does not expose the uninitialized-output contract",
752 ));
753 };
754 let mode = self
755 .dot_general_mode(entry)
756 .map_err(|error| Error::backend_source(OP, error))?;
757 cache.bind_provider_bundle(bundle_identity);
758 entry
759 .enter_or_reuse(entered, mode, |provider_context| {
760 let Some(plan) = crate::gemm::prepare_provider_gemm_into_uninit(
761 cache,
762 cache_slot,
763 lhs,
764 rhs,
765 output_shape,
766 config,
767 )?
768 else {
769 return Ok(CpuProviderOutcome::Unsupported(
773 CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Output),
774 ));
775 };
776 execute_gemm_plan_into_uninit(
777 witness,
778 provider_context,
779 plan,
780 lhs,
781 rhs,
782 accumulation,
783 output_bytes,
784 )
785 })
786 .map_err(|error| Error::backend_source(OP, error))?
787 }
788
789 #[allow(clippy::redundant_closure)]
790 fn execute_grouped(
791 &self,
792 entry: &CpuOperationEntry<'_>,
793 entered: Option<&CpuExecutionContext<'_>>,
794 lhs: TensorRead<'_>,
795 rhs: TensorRead<'_>,
796 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
797 mut output: TensorWrite<'_>,
798 ) -> Result<()> {
799 tenferro_tensor::backend::validate_grouped_gemm(
800 &lhs,
801 &rhs,
802 &output,
803 config,
804 "grouped_gemm",
805 )?;
806 if entered.is_none()
807 && self.grouped_scheduling == GroupedGemmScheduling::EngineOuter
808 && entry.supports_outer()
809 && config.jobs().len() > 1
810 {
811 if !self
812 .gemm_capabilities
813 .accepts_mode(crate::ParallelMode::Outer)
814 {
815 return Err(Error::backend_source(
816 "grouped_gemm",
817 crate::CpuProviderDomainError::ParallelModeNotSupported {
818 mode: crate::ParallelMode::Outer,
819 },
820 ));
821 }
822 return match &mut output {
823 TensorWrite::Tensor(Tensor::F32(output)) => execute_grouped_outer_typed(
824 self.gemm.as_ref(),
825 entry,
826 &lhs,
827 &rhs,
828 config,
829 output.host_data_mut()?,
830 0,
831 |view| TensorViewMut::F32(view),
832 ),
833 TensorWrite::Tensor(Tensor::F64(output)) => execute_grouped_outer_typed(
834 self.gemm.as_ref(),
835 entry,
836 &lhs,
837 &rhs,
838 config,
839 output.host_data_mut()?,
840 0,
841 |view| TensorViewMut::F64(view),
842 ),
843 TensorWrite::Tensor(Tensor::C32(output)) => execute_grouped_outer_typed(
844 self.gemm.as_ref(),
845 entry,
846 &lhs,
847 &rhs,
848 config,
849 output.host_data_mut()?,
850 0,
851 |view| TensorViewMut::C32(view),
852 ),
853 TensorWrite::Tensor(Tensor::C64(output)) => execute_grouped_outer_typed(
854 self.gemm.as_ref(),
855 entry,
856 &lhs,
857 &rhs,
858 config,
859 output.host_data_mut()?,
860 0,
861 |view| TensorViewMut::C64(view),
862 ),
863 TensorWrite::View(TensorViewMut::F32(output)) => {
864 let base = output.offset();
865 execute_grouped_outer_typed(
866 self.gemm.as_ref(),
867 entry,
868 &lhs,
869 &rhs,
870 config,
871 output.host_storage_mut()?,
872 base,
873 |view| TensorViewMut::F32(view),
874 )
875 }
876 TensorWrite::View(TensorViewMut::F64(output)) => {
877 let base = output.offset();
878 execute_grouped_outer_typed(
879 self.gemm.as_ref(),
880 entry,
881 &lhs,
882 &rhs,
883 config,
884 output.host_storage_mut()?,
885 base,
886 |view| TensorViewMut::F64(view),
887 )
888 }
889 TensorWrite::View(TensorViewMut::C32(output)) => {
890 let base = output.offset();
891 execute_grouped_outer_typed(
892 self.gemm.as_ref(),
893 entry,
894 &lhs,
895 &rhs,
896 config,
897 output.host_storage_mut()?,
898 base,
899 |view| TensorViewMut::C32(view),
900 )
901 }
902 TensorWrite::View(TensorViewMut::C64(output)) => {
903 let base = output.offset();
904 execute_grouped_outer_typed(
905 self.gemm.as_ref(),
906 entry,
907 &lhs,
908 &rhs,
909 config,
910 output.host_storage_mut()?,
911 base,
912 |view| TensorViewMut::C64(view),
913 )
914 }
915 _ => Err(unsupported_provider_error(
916 "grouped-GEMM",
917 CpuProviderUnsupported::DType(output.dtype()),
918 )),
919 };
920 }
921 let mode = self
922 .grouped_mode(entry)
923 .map_err(|error| Error::backend_source("grouped_gemm", error))?;
924 entry
925 .enter_or_reuse(entered, mode, |provider_context| {
926 let request = CpuGroupedGemmRequest::new(
927 &lhs,
928 &rhs,
929 &mut output,
930 config.jobs(),
931 config.accumulation(),
932 );
933 match self.gemm.grouped_gemm(provider_context, request)? {
934 CpuProviderOutcome::Executed => Ok(()),
935 CpuProviderOutcome::Unsupported(reason) => {
936 Err(unsupported_provider_error("grouped-GEMM", reason))
937 }
938 }
939 })
940 .map_err(|error| Error::backend_source("grouped_gemm", error))?
941 }
942}
943
944fn execute_gemm_plan(
945 provider: &dyn CpuGemmProvider,
946 context: &CpuExecutionContext<'_>,
947 plan: crate::gemm::ProviderGemmPlan,
948 lhs: &TensorRead<'_>,
949 rhs: &TensorRead<'_>,
950 accumulation: DotGeneralAccumulation,
951 output: &mut TensorWrite<'_>,
952) -> Result<CpuProviderOutcome> {
953 let batch_count = plan.batch_count();
954 let request = plan.request(lhs, rhs, output, accumulation);
955 let outcome = if batch_count == 1 {
956 provider.gemm(context, request)?
957 } else {
958 provider.strided_batched_gemm(context, request)?
959 };
960 Ok(outcome)
961}
962
963fn execute_gemm_plan_into_uninit(
964 witness: &dyn CpuUninitGemmProvider,
965 context: &CpuExecutionContext<'_>,
966 plan: crate::gemm::ProviderGemmPlan,
967 lhs: &TensorRead<'_>,
968 rhs: &TensorRead<'_>,
969 accumulation: DotGeneralAccumulation,
970 output_bytes: &mut [MaybeUninit<u8>],
971) -> Result<CpuProviderOutcome> {
972 let request = plan.uninit_request(lhs, rhs, accumulation);
973 unsafe { witness.gemm_into_uninit(context, request, output_bytes) }
978}
979
980fn canonical_gemm_fallback_supported(reason: CpuProviderUnsupported) -> bool {
981 matches!(
982 reason,
983 CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Lhs)
984 | CpuProviderUnsupported::Layout(crate::provider::CpuOperand::Rhs)
985 | CpuProviderUnsupported::Conjugation
986 )
987}
988
989fn transposed_read_view<'input>(
990 input: &TensorRead<'input>,
991 permutation: &[usize],
992) -> Result<TensorView<'input>> {
993 Ok(match input.clone().tensor_view() {
994 TensorView::F32(view) => TensorView::F32(view.transpose_view(permutation)?),
995 TensorView::F64(view) => TensorView::F64(view.transpose_view(permutation)?),
996 TensorView::I32(view) => TensorView::I32(view.transpose_view(permutation)?),
997 TensorView::I64(view) => TensorView::I64(view.transpose_view(permutation)?),
998 TensorView::Bool(view) => TensorView::Bool(view.transpose_view(permutation)?),
999 TensorView::C32(view) => TensorView::C32(view.transpose_view(permutation)?),
1000 TensorView::C64(view) => TensorView::C64(view.transpose_view(permutation)?),
1001 })
1002}
1003
1004fn pooled_zero_tensor<T>(buffers: &mut BufferPool, shape: Vec<usize>) -> Result<TypedTensor<T>>
1005where
1006 T: PoolScalar + Clone + 'static,
1007{
1008 let element_count =
1009 tenferro_tensor::validate::checked_shape_product(OP, "canonical operand", &shape)?;
1010 TypedTensor::from_vec_col_major(shape, T::pool_acquire_zeroed(buffers, element_count))
1011}
1012
1013fn allocate_canonical_operand(
1014 buffers: &mut BufferPool,
1015 dtype: DType,
1016 shape: Vec<usize>,
1017) -> Result<Tensor> {
1018 match dtype {
1019 DType::F32 => pooled_zero_tensor(buffers, shape).map(Tensor::F32),
1020 DType::F64 => pooled_zero_tensor(buffers, shape).map(Tensor::F64),
1021 DType::C32 => pooled_zero_tensor(buffers, shape).map(Tensor::C32),
1022 DType::C64 => pooled_zero_tensor(buffers, shape).map(Tensor::C64),
1023 dtype => Err(Error::unsupported_dtype(
1024 OP,
1025 dtype,
1026 crate::cpu_contraction_unsupported_dtype_message(dtype),
1027 )),
1028 }
1029}
1030
1031fn reclaim_temporary(buffers: &mut BufferPool, tensor: Tensor) {
1032 match tensor {
1033 Tensor::F32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1034 Tensor::F64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1035 Tensor::I32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1036 Tensor::I64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1037 Tensor::Bool(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1038 Tensor::C32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1039 Tensor::C64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
1040 }
1041}
1042
1043fn materialize_canonical_operand(
1044 provider: &dyn CpuLayoutTransformProvider,
1045 context: &CpuExecutionContext<'_>,
1046 buffers: &mut BufferPool,
1047 input: &TensorRead<'_>,
1048 permutation: &[usize],
1049 conjugate: bool,
1050) -> Result<Tensor> {
1051 let input_view = transposed_read_view(input, permutation)?;
1052 let dtype = input_view.dtype();
1053 let shape = input_view.shape().to_vec();
1054 let input = TensorRead::from_view(input_view);
1055 if let Some(witness) = provider.uninit_provider() {
1056 let mut output = UninitTensor::acquire(buffers, dtype, shape.clone())?;
1057 let outcome = {
1058 let output_bytes = output.as_uninit_bytes_mut();
1059 unsafe {
1064 witness.materialize_into_uninit(
1065 context,
1066 &input,
1067 CpuLayoutTransformIntent::CanonicalColumnMajor,
1068 conjugate,
1069 output_bytes,
1070 )
1071 }
1072 };
1073 match outcome {
1074 Ok(CpuProviderOutcome::Executed) => {
1075 return unsafe { output.assume_init() };
1078 }
1079 Ok(CpuProviderOutcome::Unsupported(_)) => {
1080 }
1083 Err(error) => return Err(error),
1084 }
1085 }
1086 materialize_canonical_operand_zeroed(provider, context, buffers, &input, shape, conjugate)
1087}
1088
1089fn materialize_canonical_operand_zeroed(
1090 provider: &dyn CpuLayoutTransformProvider,
1091 context: &CpuExecutionContext<'_>,
1092 buffers: &mut BufferPool,
1093 input: &TensorRead<'_>,
1094 shape: Vec<usize>,
1095 conjugate: bool,
1096) -> Result<Tensor> {
1097 let mut output = allocate_canonical_operand(buffers, input.dtype(), shape)?;
1098 let outcome = {
1099 let mut output_write = TensorWrite::from_tensor(&mut output);
1100 let request = CpuLayoutTransformRequest::new(
1101 input,
1102 &mut output_write,
1103 CpuLayoutTransformIntent::CanonicalColumnMajor,
1104 conjugate,
1105 );
1106 provider.materialize(context, request)
1107 };
1108 match outcome {
1109 Ok(CpuProviderOutcome::Executed) => Ok(output),
1110 Ok(CpuProviderOutcome::Unsupported(reason)) => {
1111 reclaim_temporary(buffers, output);
1112 Err(unsupported_provider_error("layout-transform", reason))
1113 }
1114 Err(error) => {
1115 reclaim_temporary(buffers, output);
1116 Err(error)
1117 }
1118 }
1119}
1120
1121pub(crate) enum UninitTensor<'pool> {
1128 F32(PooledUninitOutput<'pool, f32>),
1129 F64(PooledUninitOutput<'pool, f64>),
1130 C32(PooledUninitOutput<'pool, Complex32>),
1131 C64(PooledUninitOutput<'pool, Complex64>),
1132}
1133
1134impl<'pool> UninitTensor<'pool> {
1135 pub(crate) fn acquire(
1136 buffers: &'pool mut BufferPool,
1137 dtype: DType,
1138 shape: Vec<usize>,
1139 ) -> Result<Self> {
1140 match dtype {
1141 DType::F32 => Ok(Self::F32(PooledUninitOutput::new(buffers, shape)?)),
1142 DType::F64 => Ok(Self::F64(PooledUninitOutput::new(buffers, shape)?)),
1143 DType::C32 => Ok(Self::C32(PooledUninitOutput::new(buffers, shape)?)),
1144 DType::C64 => Ok(Self::C64(PooledUninitOutput::new(buffers, shape)?)),
1145 dtype => Err(Error::unsupported_dtype(
1146 OP,
1147 dtype,
1148 crate::cpu_contraction_unsupported_dtype_message(dtype),
1149 )),
1150 }
1151 }
1152
1153 pub(crate) fn as_uninit_bytes_mut(&mut self) -> &mut [MaybeUninit<u8>] {
1154 match self {
1155 Self::F32(output) => output.as_uninit_bytes_mut(),
1156 Self::F64(output) => output.as_uninit_bytes_mut(),
1157 Self::C32(output) => output.as_uninit_bytes_mut(),
1158 Self::C64(output) => output.as_uninit_bytes_mut(),
1159 }
1160 }
1161
1162 pub(crate) unsafe fn assume_init(self) -> Result<Tensor> {
1168 unsafe {
1171 match self {
1172 Self::F32(output) => output.assume_init().map(Tensor::F32),
1173 Self::F64(output) => output.assume_init().map(Tensor::F64),
1174 Self::C32(output) => output.assume_init().map(Tensor::C32),
1175 Self::C64(output) => output.assume_init().map(Tensor::C64),
1176 }
1177 }
1178 }
1179}
1180
1181fn checked_grouped_output_range(
1182 output_base: usize,
1183 output_len: usize,
1184 job: &tenferro_tensor::backend::GroupedGemmJob,
1185) -> Result<std::ops::Range<usize>> {
1186 let len = job.rows().checked_mul(job.cols()).ok_or_else(|| {
1187 Error::invalid_argument(
1188 "grouped_gemm",
1189 "jobs",
1190 "grouped-GEMM output span overflows usize",
1191 )
1192 })?;
1193 let start = output_base.checked_add(job.out_offset()).ok_or_else(|| {
1194 Error::invalid_argument(
1195 "grouped_gemm",
1196 "jobs",
1197 "grouped-GEMM output offset overflows usize",
1198 )
1199 })?;
1200 let end = start.checked_add(len).ok_or_else(|| {
1201 Error::invalid_argument(
1202 "grouped_gemm",
1203 "jobs",
1204 "grouped-GEMM output end overflows usize",
1205 )
1206 })?;
1207 if end > output_len {
1208 return Err(Error::invalid_argument(
1209 "grouped_gemm",
1210 "jobs",
1211 "grouped-GEMM output range exceeds host storage",
1212 ));
1213 }
1214 Ok(start..end)
1215}
1216
1217#[allow(clippy::too_many_arguments)]
1220fn execute_grouped_outer_typed<T>(
1221 provider: &dyn CpuGemmProvider,
1222 entry: &CpuOperationEntry<'_>,
1223 lhs: &TensorRead<'_>,
1224 rhs: &TensorRead<'_>,
1225 config: &tenferro_tensor::backend::GroupedGemmConfig<'_>,
1226 output_storage: &mut [T],
1227 output_base: isize,
1228 wrap_output: for<'a> fn(tenferro_tensor::TypedTensorViewMut<'a, T>) -> TensorViewMut<'a>,
1229) -> Result<()>
1230where
1231 T: Send + Sync + 'static,
1232{
1233 const NO_DUPLICATE: usize = usize::MAX;
1234
1235 let output_base = usize::try_from(output_base).map_err(|_| {
1236 Error::invalid_argument(
1237 "grouped_gemm",
1238 "output",
1239 "grouped-GEMM output base offset is negative",
1240 )
1241 })?;
1242 let output_storage_len = output_storage.len();
1243 for job in config.jobs() {
1244 checked_grouped_output_range(output_base, output_storage_len, job)?;
1245 }
1246
1247 let output_address = output_storage.as_mut_ptr() as usize;
1248 let operation_error = std::sync::Mutex::new(None);
1249 let job_states = PackedJobStates::new(config.jobs().len());
1250 let duplicate_index = AtomicUsize::new(NO_DUPLICATE);
1251 entry
1252 .submit_outer(config.jobs().len(), |index, provider_context| {
1253 if job_states.try_claim(index).is_err() {
1254 let _ = duplicate_index.compare_exchange(
1255 NO_DUPLICATE,
1256 index,
1257 Ordering::AcqRel,
1258 Ordering::Acquire,
1259 );
1260 return Err(CpuDomainExecutorError::Scheduling {
1261 message: format!(
1262 "executor invoked grouped-GEMM duplicate index {index}; every index must run exactly once"
1263 ),
1264 });
1265 }
1266
1267 let already_failed = operation_error
1268 .lock()
1269 .unwrap_or_else(std::sync::PoisonError::into_inner)
1270 .is_some();
1271 if !already_failed {
1272 let job = &config.jobs()[index];
1273 let result = (|| -> Result<()> {
1274 let range = checked_grouped_output_range(output_base, output_storage_len, job)?;
1275 let len = range.len();
1276 let start = range.start;
1277 let output_slice = unsafe {
1291 std::slice::from_raw_parts_mut((output_address as *mut T).add(start), len)
1292 };
1293 let output_view =
1294 tenferro_tensor::TypedTensorViewMut::from_slice([len], [1], 0, output_slice)?;
1295 let mut output = TensorWrite::from_view(wrap_output(output_view));
1296 let job = tenferro_tensor::backend::GroupedGemmJob::new(
1297 0,
1298 job.lhs_offset(),
1299 job.rhs_offset(),
1300 job.rows(),
1301 job.contracted(),
1302 job.cols(),
1303 );
1304 let request = CpuGroupedGemmRequest::new(
1305 lhs,
1306 rhs,
1307 &mut output,
1308 std::slice::from_ref(&job),
1309 config.accumulation(),
1310 );
1311 match provider.grouped_gemm(provider_context, request)? {
1312 CpuProviderOutcome::Executed => Ok(()),
1313 CpuProviderOutcome::Unsupported(reason) => {
1314 Err(unsupported_provider_error("grouped-GEMM", reason))
1315 }
1316 }
1317 })();
1318 if let Err(error) = result {
1319 *operation_error
1320 .lock()
1321 .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(error);
1322 }
1323 }
1324 let _ = job_states.complete(index);
1325 Ok(())
1326 })
1327 .map_err(|error| Error::backend_source("grouped_gemm", error))?;
1328 let duplicate = duplicate_index.load(Ordering::Acquire);
1329 if duplicate != NO_DUPLICATE {
1330 return Err(Error::backend_source(
1331 "grouped_gemm",
1332 CpuDomainExecutorError::Scheduling {
1333 message: format!(
1334 "executor invoked grouped-GEMM duplicate index {duplicate}; every index must run exactly once"
1335 ),
1336 },
1337 ));
1338 }
1339 if let Some((index, state)) = job_states.first_incomplete() {
1340 let detail = if state == GroupedJobState::Unclaimed {
1341 format!("executor omitted grouped-GEMM missing index {index}")
1342 } else {
1343 format!("executor did not complete grouped-GEMM index {index}")
1344 };
1345 return Err(Error::backend_source(
1346 "grouped_gemm",
1347 CpuDomainExecutorError::Scheduling { message: detail },
1348 ));
1349 }
1350 match operation_error.into_inner() {
1351 Ok(Some(error)) => Err(error),
1352 Err(poisoned) => poisoned.into_inner().map_or(Ok(()), Err),
1353 Ok(None) => Ok(()),
1354 }
1355}
1356
1357#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
1366#[error("missing mandatory CPU provider slots: GEMM={gemm}, layout={layout}")]
1367pub struct CpuProviderBundleBuildError {
1368 gemm: bool,
1369 layout: bool,
1370}
1371
1372#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1381pub enum CpuProviderSlot {
1382 Gemm,
1384 LayoutTransform,
1386 GeneralContraction,
1388}
1389
1390#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
1405#[non_exhaustive]
1406pub enum CpuProviderBundleInstallError {
1407 #[error(
1409 "CPU provider bundle slot {provider:?} is incompatible with domain {domain_id:?}: {source}"
1410 )]
1411 IncompatibleDomain {
1412 domain_id: tenferro_tensor::CpuDomainId,
1414 provider: CpuProviderSlot,
1416 #[source]
1418 source: CpuProviderDomainError,
1419 },
1420}
1421
1422#[derive(Debug)]
1433pub struct CpuProviderBundleBuilder {
1434 gemm: Option<Arc<dyn CpuGemmProvider>>,
1435 layout: Option<Arc<dyn CpuLayoutTransformProvider>>,
1436 general: Option<Arc<dyn CpuGeneralContractionProvider>>,
1437 general_policy: GeneralContractionPolicy,
1438 grouped_scheduling: GroupedGemmScheduling,
1439 capability_policy: ProviderCapabilityPolicy,
1440}
1441
1442impl CpuProviderBundleBuilder {
1443 pub(crate) fn provider_default_compatibility(mut self) -> Self {
1444 self.capability_policy = ProviderCapabilityPolicy::ProviderDefaultCompatibility;
1445 self
1446 }
1447
1448 pub fn gemm_provider(mut self, provider: Arc<dyn CpuGemmProvider>) -> Self {
1450 self.gemm = Some(provider);
1451 self.grouped_scheduling = GroupedGemmScheduling::ProviderOwned;
1452 self
1453 }
1454
1455 pub fn engine_outer_grouped_gemm(mut self) -> Self {
1462 self.grouped_scheduling = GroupedGemmScheduling::EngineOuter;
1463 self
1464 }
1465
1466 pub fn layout_transform_provider(
1468 mut self,
1469 provider: Arc<dyn CpuLayoutTransformProvider>,
1470 ) -> Self {
1471 self.layout = Some(provider);
1472 self
1473 }
1474
1475 pub fn prefer_general_contraction_provider(
1477 mut self,
1478 provider: Arc<dyn CpuGeneralContractionProvider>,
1479 ) -> Self {
1480 self.general = Some(provider);
1481 self.general_policy = GeneralContractionPolicy::Preferred;
1482 self
1483 }
1484
1485 pub fn require_general_contraction_provider(
1487 mut self,
1488 provider: Arc<dyn CpuGeneralContractionProvider>,
1489 ) -> Self {
1490 self.general = Some(provider);
1491 self.general_policy = GeneralContractionPolicy::Required;
1492 self
1493 }
1494
1495 pub fn build(self) -> std::result::Result<CpuProviderBundle, CpuProviderBundleBuildError> {
1501 let missing = CpuProviderBundleBuildError {
1502 gemm: self.gemm.is_none(),
1503 layout: self.layout.is_none(),
1504 };
1505 let (Some(gemm), Some(layout)) = (self.gemm, self.layout) else {
1506 return Err(missing);
1507 };
1508 let general_capabilities = self
1509 .general
1510 .as_ref()
1511 .map(|provider| provider.execution_capabilities());
1512 let gemm_capabilities = gemm.execution_capabilities();
1513 let layout_capabilities = layout.execution_capabilities();
1514 Ok(CpuProviderBundle {
1515 inner: Arc::new(CpuProviderBundleInner {
1516 dot_general: DotGeneralRuntime {
1517 general: self.general,
1518 gemm,
1519 layout,
1520 general_capabilities,
1521 gemm_capabilities,
1522 layout_capabilities,
1523 general_policy: self.general_policy,
1524 grouped_scheduling: self.grouped_scheduling,
1525 capability_policy: self.capability_policy,
1526 },
1527 }),
1528 })
1529 }
1530}
1531
1532fn validate_axis_ranges(axes: &[usize], rank: usize) -> Result<()> {
1533 for &axis in axes {
1534 if axis >= rank {
1535 return Err(Error::axis_out_of_bounds(OP, axis, rank));
1536 }
1537 }
1538 Ok(())
1539}
1540
1541fn role_mask(axes: &[usize], rank: usize, role: &'static str) -> Result<Option<u64>> {
1542 if rank > 64 {
1543 for (position, &axis) in axes.iter().enumerate() {
1544 if axes[..position].contains(&axis) {
1545 return Err(Error::duplicate_axis(OP, axis, role));
1546 }
1547 }
1548 return Ok(None);
1549 }
1550
1551 let mut mask = 0_u64;
1552 for &axis in axes {
1553 let bit = 1_u64 << axis;
1554 if mask & bit != 0 {
1555 return Err(Error::duplicate_axis(OP, axis, role));
1556 }
1557 mask |= bit;
1558 }
1559 Ok(Some(mask))
1560}
1561
1562fn validate_disjoint(
1563 first: &[usize],
1564 first_mask: Option<u64>,
1565 first_role: &'static str,
1566 second: &[usize],
1567 second_mask: Option<u64>,
1568 second_role: &'static str,
1569) -> Result<()> {
1570 let overlap = match (first_mask, second_mask) {
1571 (Some(first), Some(second)) => first & second,
1572 _ => 0,
1573 };
1574 let conflict = if overlap != 0 || first_mask.is_none() {
1575 first.iter().copied().find(|axis| second.contains(axis))
1576 } else {
1577 None
1578 };
1579 if let Some(axis) = conflict {
1580 return Err(Error::validation(
1581 OP,
1582 ValidationError::AxisRoleConflict {
1583 axis,
1584 first_role,
1585 second_role,
1586 },
1587 ));
1588 }
1589 Ok(())
1590}
1591
1592pub(crate) fn validate_axis_groups<'a>(
1593 lhs_rank: usize,
1594 rhs_rank: usize,
1595 config: &'a DotGeneralConfig,
1596) -> Result<CpuContractionAxes<'a>> {
1597 validate_axis_ranges(&config.lhs_contracting_dims, lhs_rank)?;
1598 validate_axis_ranges(&config.rhs_contracting_dims, rhs_rank)?;
1599 validate_axis_ranges(&config.lhs_batch_dims, lhs_rank)?;
1600 validate_axis_ranges(&config.rhs_batch_dims, rhs_rank)?;
1601
1602 let lhs_contracting_mask = role_mask(
1603 &config.lhs_contracting_dims,
1604 lhs_rank,
1605 "lhs_contracting_dims",
1606 )?;
1607 let rhs_contracting_mask = role_mask(
1608 &config.rhs_contracting_dims,
1609 rhs_rank,
1610 "rhs_contracting_dims",
1611 )?;
1612 let lhs_batch_mask = role_mask(&config.lhs_batch_dims, lhs_rank, "lhs_batch_dims")?;
1613 let rhs_batch_mask = role_mask(&config.rhs_batch_dims, rhs_rank, "rhs_batch_dims")?;
1614
1615 validate_disjoint(
1616 &config.lhs_contracting_dims,
1617 lhs_contracting_mask,
1618 "lhs contracting",
1619 &config.lhs_batch_dims,
1620 lhs_batch_mask,
1621 "lhs batch",
1622 )?;
1623 validate_disjoint(
1624 &config.rhs_contracting_dims,
1625 rhs_contracting_mask,
1626 "rhs contracting",
1627 &config.rhs_batch_dims,
1628 rhs_batch_mask,
1629 "rhs batch",
1630 )?;
1631
1632 if config.lhs_contracting_dims.len() != config.rhs_contracting_dims.len() {
1633 return Err(Error::invalid_argument(
1634 OP,
1635 "dot_general_config",
1636 format!(
1637 "lhs/rhs contracting dim counts differ ({} vs {})",
1638 config.lhs_contracting_dims.len(),
1639 config.rhs_contracting_dims.len(),
1640 ),
1641 ));
1642 }
1643 if config.lhs_batch_dims.len() != config.rhs_batch_dims.len() {
1644 return Err(Error::invalid_argument(
1645 OP,
1646 "dot_general_config",
1647 format!(
1648 "lhs/rhs batch dim counts differ ({} vs {})",
1649 config.lhs_batch_dims.len(),
1650 config.rhs_batch_dims.len(),
1651 ),
1652 ));
1653 }
1654
1655 Ok(CpuContractionAxes::new(
1656 lhs_rank,
1657 rhs_rank,
1658 &config.lhs_contracting_dims,
1659 &config.rhs_contracting_dims,
1660 &config.lhs_batch_dims,
1661 &config.rhs_batch_dims,
1662 lhs_contracting_mask.zip(lhs_batch_mask).map(|(a, b)| a | b),
1663 rhs_contracting_mask.zip(rhs_batch_mask).map(|(a, b)| a | b),
1664 ))
1665}
1666
1667#[derive(Clone, Copy, Debug)]
1668pub(crate) struct ValidatedDotGeneral<'a> {
1669 axes: CpuContractionAxes<'a>,
1670 output_element_count: usize,
1671}
1672
1673impl<'a> ValidatedDotGeneral<'a> {
1674 pub(crate) fn axes(&self) -> &CpuContractionAxes<'a> {
1675 &self.axes
1676 }
1677
1678 pub(crate) fn output_element_count(&self) -> usize {
1679 self.output_element_count
1680 }
1681
1682 #[allow(dead_code)]
1683 pub(crate) fn request<'request, 'input, 'output>(
1684 &'request self,
1685 lhs: &'request TensorRead<'input>,
1686 rhs: &'request TensorRead<'input>,
1687 output: &'request mut TensorWrite<'output>,
1688 accumulation: DotGeneralAccumulation,
1689 ) -> CpuDotGeneralRequest<'request, 'input, 'output>
1690 where
1691 'a: 'request,
1692 {
1693 CpuDotGeneralRequest::new(lhs, rhs, output, self.axes, accumulation)
1694 }
1695}
1696
1697fn validate_paired_extents(
1698 lhs: &TensorRead<'_>,
1699 rhs: &TensorRead<'_>,
1700 axes: &CpuContractionAxes<'_>,
1701) -> Result<()> {
1702 for (lhs_axis, rhs_axis) in axes.contracting_pairs().chain(axes.batch_pairs()) {
1703 if lhs.shape()[lhs_axis] != rhs.shape()[rhs_axis] {
1704 return Err(Error::validation(
1705 OP,
1706 ShapeMismatch::ContractedDimensions {
1707 lhs_axis,
1708 lhs_size: lhs.shape()[lhs_axis],
1709 rhs_axis,
1710 rhs_size: rhs.shape()[rhs_axis],
1711 }
1712 .into(),
1713 ));
1714 }
1715 }
1716 Ok(())
1717}
1718
1719fn expected_output_shape(
1720 lhs: &TensorRead<'_>,
1721 rhs: &TensorRead<'_>,
1722 axes: &CpuContractionAxes<'_>,
1723) -> Vec<usize> {
1724 axes.lhs_free_axes()
1725 .map(|axis| lhs.shape()[axis])
1726 .chain(axes.rhs_free_axes().map(|axis| rhs.shape()[axis]))
1727 .chain(
1728 axes.batch_pairs()
1729 .map(|(lhs_axis, _)| lhs.shape()[lhs_axis]),
1730 )
1731 .collect()
1732}
1733
1734fn output_shape_matches(
1735 lhs: &TensorRead<'_>,
1736 rhs: &TensorRead<'_>,
1737 output: &TensorWrite<'_>,
1738 axes: &CpuContractionAxes<'_>,
1739) -> Result<()> {
1740 let expected_rank =
1741 axes.lhs_free_axes().count() + axes.rhs_free_axes().count() + axes.batch_pairs().len();
1742 let mut actual = output.shape().iter().copied();
1743 let matches = output.shape().len() == expected_rank
1744 && axes
1745 .lhs_free_axes()
1746 .map(|axis| lhs.shape()[axis])
1747 .chain(axes.rhs_free_axes().map(|axis| rhs.shape()[axis]))
1748 .chain(
1749 axes.batch_pairs()
1750 .map(|(lhs_axis, _)| lhs.shape()[lhs_axis]),
1751 )
1752 .all(|expected| actual.next() == Some(expected));
1753 if matches {
1754 return Ok(());
1755 }
1756
1757 Err(Error::validation(
1758 OP,
1759 ShapeMismatch::ExpectedActual {
1760 expected: expected_output_shape(lhs, rhs, axes).into(),
1761 actual: output.shape().to_vec().into(),
1762 }
1763 .into(),
1764 ))
1765}
1766
1767fn layout_overflow() -> Error {
1768 Error::validation(OP, ValidationError::IntegerOverflow)
1769}
1770
1771pub(crate) fn validate_layout_metadata(
1772 role: &'static str,
1773 shape: &[usize],
1774 strides: &[isize],
1775 offset: isize,
1776 storage_len: usize,
1777) -> Result<usize> {
1778 if shape.len() != strides.len() {
1779 return Err(Error::validation(
1780 OP,
1781 ValidationError::RankMismatch {
1782 expected: shape.len(),
1783 actual: strides.len(),
1784 },
1785 ));
1786 }
1787 let element_count = tenferro_tensor::validate::checked_shape_product(OP, role, shape)?;
1788
1789 if shape.contains(&0) {
1790 let offset = usize::try_from(offset).map_err(|_| {
1791 Error::invalid_argument(OP, role, "minimum reachable offset is negative")
1792 })?;
1793 if offset > storage_len {
1794 return Err(Error::validation(OP, ValidationError::ViewOutOfBounds));
1795 }
1796 return Ok(element_count);
1797 }
1798
1799 let mut minimum = offset;
1800 let mut maximum = offset;
1801 for (&extent, &stride) in shape.iter().zip(strides) {
1802 let steps = isize::try_from(extent - 1).map_err(|_| layout_overflow())?;
1803 let end = stride.checked_mul(steps).ok_or_else(layout_overflow)?;
1804 let (axis_minimum, axis_maximum) = if end < 0 { (end, 0) } else { (0, end) };
1805 minimum = minimum
1806 .checked_add(axis_minimum)
1807 .ok_or_else(layout_overflow)?;
1808 maximum = maximum
1809 .checked_add(axis_maximum)
1810 .ok_or_else(layout_overflow)?;
1811 }
1812 let minimum = usize::try_from(minimum)
1813 .map_err(|_| Error::invalid_argument(OP, role, "minimum reachable offset is negative"))?;
1814 let maximum = usize::try_from(maximum)
1815 .map_err(|_| Error::invalid_argument(OP, role, "maximum reachable offset is negative"))?;
1816 if minimum > maximum || maximum >= storage_len {
1817 return Err(Error::validation(OP, ValidationError::ViewOutOfBounds));
1818 }
1819 Ok(element_count)
1820}
1821
1822macro_rules! validate_owned_layout {
1823 ($tensor:expr, $role:expr) => {{
1824 let tensor = $tensor;
1825 if tensor.backend_buffer().is_some() {
1826 return Err(crate::cpu_backend_buffer_error(OP));
1827 }
1828 let storage_len = tensor.host_data()?.len();
1829 validate_layout_metadata(
1830 $role,
1831 tensor.shape(),
1832 tensor.layout().strides(),
1833 tensor.layout().offset(),
1834 storage_len,
1835 )
1836 }};
1837}
1838
1839macro_rules! validate_read_view_layout {
1840 ($view:expr, $role:expr) => {{
1841 let view = $view;
1842 let storage_len = view.host_storage()?.len();
1843 validate_layout_metadata(
1844 $role,
1845 view.shape(),
1846 view.strides(),
1847 view.offset(),
1848 storage_len,
1849 )
1850 }};
1851}
1852
1853macro_rules! validate_write_view_layout {
1854 ($view:expr, $role:expr) => {{
1855 let view = $view;
1856 let storage_len = view.host_storage()?.len();
1857 validate_layout_metadata(
1858 $role,
1859 view.shape(),
1860 view.strides(),
1861 view.offset(),
1862 storage_len,
1863 )
1864 }};
1865}
1866
1867fn validate_read_layout(tensor: &TensorRead<'_>, role: &'static str) -> Result<usize> {
1868 match tensor {
1869 TensorRead::Tensor(tensor) => match tensor {
1870 Tensor::F32(tensor) => validate_owned_layout!(tensor, role),
1871 Tensor::F64(tensor) => validate_owned_layout!(tensor, role),
1872 Tensor::I32(tensor) => validate_owned_layout!(tensor, role),
1873 Tensor::I64(tensor) => validate_owned_layout!(tensor, role),
1874 Tensor::Bool(tensor) => validate_owned_layout!(tensor, role),
1875 Tensor::C32(tensor) => validate_owned_layout!(tensor, role),
1876 Tensor::C64(tensor) => validate_owned_layout!(tensor, role),
1877 },
1878 TensorRead::View(view) => match view {
1879 TensorView::F32(view) => validate_read_view_layout!(view, role),
1880 TensorView::F64(view) => validate_read_view_layout!(view, role),
1881 TensorView::I32(view) => validate_read_view_layout!(view, role),
1882 TensorView::I64(view) => validate_read_view_layout!(view, role),
1883 TensorView::Bool(view) => validate_read_view_layout!(view, role),
1884 TensorView::C32(view) => validate_read_view_layout!(view, role),
1885 TensorView::C64(view) => validate_read_view_layout!(view, role),
1886 },
1887 }
1888}
1889
1890fn validate_write_layout(tensor: &TensorWrite<'_>, role: &'static str) -> Result<usize> {
1891 match tensor {
1892 TensorWrite::Tensor(tensor) => match tensor {
1893 Tensor::F32(tensor) => validate_owned_layout!(tensor, role),
1894 Tensor::F64(tensor) => validate_owned_layout!(tensor, role),
1895 Tensor::I32(tensor) => validate_owned_layout!(tensor, role),
1896 Tensor::I64(tensor) => validate_owned_layout!(tensor, role),
1897 Tensor::Bool(tensor) => validate_owned_layout!(tensor, role),
1898 Tensor::C32(tensor) => validate_owned_layout!(tensor, role),
1899 Tensor::C64(tensor) => validate_owned_layout!(tensor, role),
1900 },
1901 TensorWrite::View(view) => match view {
1902 TensorViewMut::F32(view) => validate_write_view_layout!(view, role),
1903 TensorViewMut::F64(view) => validate_write_view_layout!(view, role),
1904 TensorViewMut::I32(view) => validate_write_view_layout!(view, role),
1905 TensorViewMut::I64(view) => validate_write_view_layout!(view, role),
1906 TensorViewMut::Bool(view) => validate_write_view_layout!(view, role),
1907 TensorViewMut::C32(view) => validate_write_view_layout!(view, role),
1908 TensorViewMut::C64(view) => validate_write_view_layout!(view, role),
1909 },
1910 }
1911}
1912
1913pub(crate) fn validate_dot_general<'a>(
1914 lhs: &TensorRead<'_>,
1915 rhs: &TensorRead<'_>,
1916 output: &TensorWrite<'_>,
1917 config: &'a DotGeneralConfig,
1918 accumulation: DotGeneralAccumulation,
1919) -> Result<ValidatedDotGeneral<'a>> {
1920 if lhs.dtype() != rhs.dtype() {
1921 return Err(Error::dtype_mismatch(OP, lhs.dtype(), rhs.dtype()));
1922 }
1923 if output.dtype() != lhs.dtype() {
1924 return Err(Error::dtype_mismatch(OP, output.dtype(), lhs.dtype()));
1925 }
1926 if accumulation.alpha.dtype() != lhs.dtype() {
1927 return Err(Error::dtype_mismatch(
1928 OP,
1929 lhs.dtype(),
1930 accumulation.alpha.dtype(),
1931 ));
1932 }
1933 if accumulation.beta.dtype() != lhs.dtype() {
1934 return Err(Error::dtype_mismatch(
1935 OP,
1936 lhs.dtype(),
1937 accumulation.beta.dtype(),
1938 ));
1939 }
1940
1941 crate::structural::validate_cpu_host_placement(OP, "lhs", read_placement(lhs))?;
1942 crate::structural::validate_cpu_host_placement(OP, "rhs", read_placement(rhs))?;
1943 crate::structural::validate_cpu_host_placement(OP, "output", write_placement(output))?;
1944 validate_read_layout(lhs, "lhs")?;
1945 validate_read_layout(rhs, "rhs")?;
1946 let output_element_count = validate_write_layout(output, "output")?;
1947
1948 let axes = validate_axis_groups(lhs.shape().len(), rhs.shape().len(), config)?;
1949 validate_paired_extents(lhs, rhs, &axes)?;
1950 output_shape_matches(lhs, rhs, output, &axes)?;
1951
1952 Ok(ValidatedDotGeneral {
1953 axes,
1954 output_element_count,
1955 })
1956}
1957
1958fn read_placement<'a>(tensor: &'a TensorRead<'_>) -> &'a tenferro_tensor::Placement {
1959 match tensor {
1960 TensorRead::Tensor(tensor) => tensor.placement(),
1961 TensorRead::View(view) => match view {
1962 tenferro_tensor::TensorView::F32(view) => view.placement(),
1963 tenferro_tensor::TensorView::F64(view) => view.placement(),
1964 tenferro_tensor::TensorView::I32(view) => view.placement(),
1965 tenferro_tensor::TensorView::I64(view) => view.placement(),
1966 tenferro_tensor::TensorView::Bool(view) => view.placement(),
1967 tenferro_tensor::TensorView::C32(view) => view.placement(),
1968 tenferro_tensor::TensorView::C64(view) => view.placement(),
1969 },
1970 }
1971}
1972
1973fn write_placement<'a>(tensor: &'a TensorWrite<'_>) -> &'a tenferro_tensor::Placement {
1974 match tensor {
1975 TensorWrite::Tensor(tensor) => tensor.placement(),
1976 TensorWrite::View(view) => match view {
1977 tenferro_tensor::TensorViewMut::F32(view) => view.placement(),
1978 tenferro_tensor::TensorViewMut::F64(view) => view.placement(),
1979 tenferro_tensor::TensorViewMut::I32(view) => view.placement(),
1980 tenferro_tensor::TensorViewMut::I64(view) => view.placement(),
1981 tenferro_tensor::TensorViewMut::Bool(view) => view.placement(),
1982 tenferro_tensor::TensorViewMut::C32(view) => view.placement(),
1983 tenferro_tensor::TensorViewMut::C64(view) => view.placement(),
1984 },
1985 }
1986}
1987
1988#[cfg(test)]
1989mod tests;