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