Skip to main content

tenferro_cpu/
dot_runtime.rs

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/// Policy applied when the configured general-contraction provider reports a
26/// typed capability miss.
27///
28/// # Examples
29///
30/// ```
31/// use tenferro_cpu::GeneralContractionPolicy;
32/// assert_ne!(
33///     GeneralContractionPolicy::Preferred,
34///     GeneralContractionPolicy::Required,
35/// );
36/// ```
37#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
38pub enum GeneralContractionPolicy {
39    /// Continue to the configured layout-plus-GEMM path.
40    #[default]
41    Preferred,
42    /// Convert a capability miss into a structured unsupported error.
43    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
96// INVARIANT: the public safe executor boundary can independently duplicate or
97// omit any grouped job, so sound post-submit auditing requires O(job_count)
98// state with at least UNCLAIMED/RUNNING/COMPLETE. Packing two bits per job into
99// four inline AtomicUsize words covers 2 * usize::BITS jobs without allocation;
100// only larger groups spill. Whole-word CAS updates preserve neighboring states.
101struct 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/// Immutable direct provider slots installed on a CPU backend.
200///
201/// Clones share the same slot identity and may safely share compatible
202/// analysis-cache entries.
203///
204/// # Examples
205///
206/// ```
207/// use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
208/// let bundle = CpuProviderBundle::builder(CpuBackendKind::default_compiled()).build()?;
209/// let cloned = bundle.clone();
210/// assert!(bundle.shares_identity_with(&cloned));
211/// # Ok::<(), tenferro_cpu::CpuProviderBundleBuildError>(())
212/// ```
213#[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    /// Start a bundle builder with the standard providers for `kind`.
246    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    /// Start an empty custom builder.
258    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    /// Return whether two handles share one immutable provider identity.
270    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    // INVARIANT: these arguments are distinct borrowed components of one
552    // validated dispatch; grouping them would duplicate validation-owned
553    // metadata or add a request allocation to the hot path.
554    #[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// INVARIANT: provider, context, tensor views, grouped metadata, and output
1012// storage are independent borrowed parts of one already-validated request.
1013#[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                // INVARIANT: the immutable job, output base, and allocation
1072                // length are identical to preflight. The shared checked helper
1073                // therefore reconstructs the same in-bounds range inside this
1074                // worker; the common grouped validator also proved distinct
1075                // job ranges disjoint. Before reaching this point, the
1076                // packed atomic claim changed this job from UNCLAIMED to
1077                // RUNNING without clobbering neighboring states, so even a
1078                // contract-violating safe executor cannot send a second
1079                // invocation of this index to the provider.
1080                // SAFETY: `start..start + len` is in this allocation. Distinct
1081                // jobs have disjoint validated ranges, and the atomic claim
1082                // permits exactly one invocation of each job to construct its
1083                // mutable slice.
1084                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/// Error returned when a custom CPU provider bundle omits mandatory slots.
1152///
1153/// # Examples
1154///
1155/// ```
1156/// use tenferro_cpu::CpuProviderBundle;
1157/// assert!(CpuProviderBundle::custom_builder().build().is_err());
1158/// ```
1159#[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/// Provider slot that failed construction-time domain validation.
1167///
1168/// # Examples
1169///
1170/// ```
1171/// use tenferro_cpu::CpuProviderSlot;
1172/// assert_ne!(CpuProviderSlot::Gemm, CpuProviderSlot::LayoutTransform);
1173/// ```
1174#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1175pub enum CpuProviderSlot {
1176    /// GEMM, strided-batched GEMM, and grouped-GEMM provider.
1177    Gemm,
1178    /// Layout materialization provider.
1179    LayoutTransform,
1180    /// Optional complete general-contraction provider.
1181    GeneralContraction,
1182}
1183
1184/// Failure to install a CPU provider bundle for the backend's domains.
1185///
1186/// Phase 2 reserves this typed surface for construction-time domain/provider
1187/// validation. Provider capability classification populates concrete
1188/// incompatibilities without adding a second installation API.
1189///
1190/// # Examples
1191///
1192/// ```
1193/// use tenferro_cpu::CpuProviderBundleInstallError;
1194/// # fn diagnostic(error: &CpuProviderBundleInstallError) -> String {
1195/// error.to_string()
1196/// # }
1197/// ```
1198#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
1199#[non_exhaustive]
1200pub enum CpuProviderBundleInstallError {
1201    /// A provider capability cannot satisfy one selected resource domain.
1202    #[error(
1203        "CPU provider bundle slot {provider:?} is incompatible with domain {domain_id:?}: {source}"
1204    )]
1205    IncompatibleDomain {
1206        /// Domain rejected by construction-time validation.
1207        domain_id: tenferro_tensor::CpuDomainId,
1208        /// Provider slot rejected by the domain contract.
1209        provider: CpuProviderSlot,
1210        /// Typed count, placement, or parallel-mode incompatibility.
1211        #[source]
1212        source: CpuProviderDomainError,
1213    },
1214}
1215
1216/// Construction-time builder for immutable CPU provider slots.
1217///
1218/// # Examples
1219///
1220/// ```
1221/// use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
1222/// let bundle = CpuProviderBundle::builder(CpuBackendKind::default_compiled()).build()?;
1223/// assert!(bundle.shares_identity_with(&bundle.clone()));
1224/// # Ok::<(), tenferro_cpu::CpuProviderBundleBuildError>(())
1225/// ```
1226#[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    /// Replace the GEMM-family provider slot.
1243    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    /// Permit the engine to fan out grouped GEMM into concurrent single-job calls.
1250    ///
1251    /// The installed GEMM provider must be safe for concurrent calls and must
1252    /// honor [`crate::provider::ParallelMode::Sequential`] without creating inner
1253    /// workers. Custom providers remain provider-owned unless this capability
1254    /// is selected explicitly.
1255    pub fn engine_outer_grouped_gemm(mut self) -> Self {
1256        self.grouped_scheduling = GroupedGemmScheduling::EngineOuter;
1257        self
1258    }
1259
1260    /// Replace the layout-materialization provider slot.
1261    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    /// Install a preferred general-contraction provider.
1270    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    /// Install a required general-contraction provider.
1280    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    /// Validate the mandatory slots and freeze the bundle identity.
1290    ///
1291    /// # Errors
1292    ///
1293    /// Returns [`CpuProviderBundleBuildError`] when GEMM or layout is absent.
1294    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;