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