Skip to main content

tenferro_cpu/
dot_runtime.rs

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/// Policy applied when the configured general-contraction provider reports a
28/// typed capability miss.
29///
30/// # Examples
31///
32/// ```
33/// use tenferro_cpu::GeneralContractionPolicy;
34/// assert_ne!(
35///     GeneralContractionPolicy::Preferred,
36///     GeneralContractionPolicy::Required,
37/// );
38/// ```
39#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
40pub enum GeneralContractionPolicy {
41    /// Continue to the configured layout-plus-GEMM path.
42    #[default]
43    Preferred,
44    /// Convert a capability miss into a structured unsupported error.
45    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
98// INVARIANT: the public safe executor boundary can independently duplicate or
99// omit any grouped job, so sound post-submit auditing requires O(job_count)
100// state with at least UNCLAIMED/RUNNING/COMPLETE. Packing two bits per job into
101// four inline AtomicUsize words covers 2 * usize::BITS jobs without allocation;
102// only larger groups spill. Whole-word CAS updates preserve neighboring states.
103struct 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/// Immutable direct provider slots installed on a CPU backend.
212///
213/// Clones share the same slot identity and may safely share compatible
214/// analysis-cache entries.
215///
216/// # Examples
217///
218/// ```
219/// use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
220/// let bundle = CpuProviderBundle::builder(CpuBackendKind::default_compiled()).build()?;
221/// let cloned = bundle.clone();
222/// assert!(bundle.shares_identity_with(&cloned));
223/// # Ok::<(), tenferro_cpu::CpuProviderBundleBuildError>(())
224/// ```
225#[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    /// Start a bundle builder with the standard providers for `kind`.
258    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    /// Start an empty custom builder.
270    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    /// Return whether two handles share one immutable provider identity.
282    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    // INVARIANT: these arguments are distinct borrowed components of one
574    // validated dispatch; grouping them would duplicate validation-owned
575    // metadata or add a request allocation to the hot path.
576    #[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    /// Execute a `beta == 0` allocated dot into uninitialized pooled bytes.
722    ///
723    /// Returns [`CpuProviderOutcome::Executed`] after every destination
724    /// element is initialized, or [`CpuProviderOutcome::Unsupported`] when the
725    /// GEMM provider cannot execute the planned contraction into
726    /// uninitialized storage (the caller discards the checkout and retries on
727    /// the zeroed path). Errors propagate; a provider error may follow a
728    /// partial write, so it is never silently retried.
729    ///
730    /// Only the direct GEMM plan is attempted here: the uninit checkout holds
731    /// the scratch pool exclusively, so the canonical fallback (which
732    /// materializes operands from the pool) is left to the zeroed path.
733    #[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                    // No direct plan; the canonical path needs the scratch
770                    // pool, which is exclusively held by the uninit checkout.
771                    // The caller falls back to the zeroed path.
772                    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    // SAFETY: the witness is structural proof the provider asserted the
974    // full-overwrite contract via `unsafe impl`; the caller guarantees
975    // beta == 0, so every destination element is written before `Executed`
976    // and never read.
977    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            // SAFETY: `witness` is structural proof the provider asserted the
1060            // full-overwrite contract via `unsafe impl`; `Executed` means
1061            // every element of `output_bytes` was written by
1062            // `materialize_into_uninit` (never read).
1063            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                // SAFETY: the unsafe provider contract guarantees the
1076                // destination is fully initialized before `Executed`.
1077                return unsafe { output.assume_init() };
1078            }
1079            Ok(CpuProviderOutcome::Unsupported(_)) => {
1080                // Discard the uninit checkout (drop frees via
1081                // `pool_discard_uninit`) and fall back to the zeroed path.
1082            }
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
1121/// Dtype-dispatched pooled full-overwrite destination for the uninitialized
1122/// dot paths.
1123///
1124/// The destination travels only as `MaybeUninit` bytes until an unsafe
1125/// `assume_init` completes the handoff; no `TensorWrite` is ever fabricated
1126/// over uninitialized storage.
1127pub(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    /// # Safety
1163    ///
1164    /// Every logical destination element must have been initialized by the
1165    /// completed unsafe provider call before this handoff; otherwise reading
1166    /// or dropping the returned tensor is undefined behavior.
1167    pub(crate) unsafe fn assume_init(self) -> Result<Tensor> {
1168        // SAFETY: the caller proves every logical destination element was
1169        // written before `Executed` by the unsafe provider impl.
1170        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// INVARIANT: provider, context, tensor views, grouped metadata, and output
1218// storage are independent borrowed parts of one already-validated request.
1219#[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                // INVARIANT: the immutable job, output base, and allocation
1278                // length are identical to preflight. The shared checked helper
1279                // therefore reconstructs the same in-bounds range inside this
1280                // worker; the common grouped validator also proved distinct
1281                // job ranges disjoint. Before reaching this point, the
1282                // packed atomic claim changed this job from UNCLAIMED to
1283                // RUNNING without clobbering neighboring states, so even a
1284                // contract-violating safe executor cannot send a second
1285                // invocation of this index to the provider.
1286                // SAFETY: `start..start + len` is in this allocation. Distinct
1287                // jobs have disjoint validated ranges, and the atomic claim
1288                // permits exactly one invocation of each job to construct its
1289                // mutable slice.
1290                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/// Error returned when a custom CPU provider bundle omits mandatory slots.
1358///
1359/// # Examples
1360///
1361/// ```
1362/// use tenferro_cpu::CpuProviderBundle;
1363/// assert!(CpuProviderBundle::custom_builder().build().is_err());
1364/// ```
1365#[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/// Provider slot that failed construction-time domain validation.
1373///
1374/// # Examples
1375///
1376/// ```
1377/// use tenferro_cpu::CpuProviderSlot;
1378/// assert_ne!(CpuProviderSlot::Gemm, CpuProviderSlot::LayoutTransform);
1379/// ```
1380#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1381pub enum CpuProviderSlot {
1382    /// GEMM, strided-batched GEMM, and grouped-GEMM provider.
1383    Gemm,
1384    /// Layout materialization provider.
1385    LayoutTransform,
1386    /// Optional complete general-contraction provider.
1387    GeneralContraction,
1388}
1389
1390/// Failure to install a CPU provider bundle for the backend's domains.
1391///
1392/// Phase 2 reserves this typed surface for construction-time domain/provider
1393/// validation. Provider capability classification populates concrete
1394/// incompatibilities without adding a second installation API.
1395///
1396/// # Examples
1397///
1398/// ```
1399/// use tenferro_cpu::CpuProviderBundleInstallError;
1400/// # fn diagnostic(error: &CpuProviderBundleInstallError) -> String {
1401/// error.to_string()
1402/// # }
1403/// ```
1404#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
1405#[non_exhaustive]
1406pub enum CpuProviderBundleInstallError {
1407    /// A provider capability cannot satisfy one selected resource domain.
1408    #[error(
1409        "CPU provider bundle slot {provider:?} is incompatible with domain {domain_id:?}: {source}"
1410    )]
1411    IncompatibleDomain {
1412        /// Domain rejected by construction-time validation.
1413        domain_id: tenferro_tensor::CpuDomainId,
1414        /// Provider slot rejected by the domain contract.
1415        provider: CpuProviderSlot,
1416        /// Typed count, placement, or parallel-mode incompatibility.
1417        #[source]
1418        source: CpuProviderDomainError,
1419    },
1420}
1421
1422/// Construction-time builder for immutable CPU provider slots.
1423///
1424/// # Examples
1425///
1426/// ```
1427/// use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
1428/// let bundle = CpuProviderBundle::builder(CpuBackendKind::default_compiled()).build()?;
1429/// assert!(bundle.shares_identity_with(&bundle.clone()));
1430/// # Ok::<(), tenferro_cpu::CpuProviderBundleBuildError>(())
1431/// ```
1432#[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    /// Replace the GEMM-family provider slot.
1449    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    /// Permit the engine to fan out grouped GEMM into concurrent single-job calls.
1456    ///
1457    /// The installed GEMM provider must be safe for concurrent calls and must
1458    /// honor [`crate::provider::ParallelMode::Sequential`] without creating inner
1459    /// workers. Custom providers remain provider-owned unless this capability
1460    /// is selected explicitly.
1461    pub fn engine_outer_grouped_gemm(mut self) -> Self {
1462        self.grouped_scheduling = GroupedGemmScheduling::EngineOuter;
1463        self
1464    }
1465
1466    /// Replace the layout-materialization provider slot.
1467    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    /// Install a preferred general-contraction provider.
1476    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    /// Install a required general-contraction provider.
1486    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    /// Validate the mandatory slots and freeze the bundle identity.
1496    ///
1497    /// # Errors
1498    ///
1499    /// Returns [`CpuProviderBundleBuildError`] when GEMM or layout is absent.
1500    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;