Skip to main content

tenferro_cpu/
provider.rs

1//! Object-safe CPU contraction provider contracts.
2//!
3//! Providers synchronously write into engine-owned outputs. Request
4//! constructors are crate-private because only the CPU engine may attest that
5//! tensor metadata and reachable ranges have already been validated.
6
7use core::fmt;
8use std::mem::MaybeUninit;
9use std::num::NonZeroUsize;
10use std::sync::Arc;
11
12use tenferro_tensor::backend::GroupedGemmJob;
13use tenferro_tensor::{
14    DType, DotGeneralAccumulation, Tensor, TensorRead, TensorView, TensorViewMut, TensorWrite,
15};
16
17use crate::arbiter::{with_execution_owner, ResourcePermit};
18use crate::backend::CpuBackendKind;
19use crate::buffer_pool::BufferPool;
20use crate::domain_executor::{indexed_jobs, install_scoped};
21/// The Rust scalar type behind a preset variant name a macro received.
22macro_rules! preset_scalar {
23    (F32) => {
24        f32
25    };
26    (F64) => {
27        f64
28    };
29    (I32) => {
30        i32
31    };
32    (I64) => {
33        i64
34    };
35    (Bool) => {
36        bool
37    };
38    (C32) => {
39        num_complex::Complex32
40    };
41    (C64) => {
42        num_complex::Complex64
43    };
44}
45
46#[cfg(feature = "cpu-blas")]
47use crate::provider_capability::builtin_blas_execution_capabilities;
48#[cfg(not(feature = "cpu-blas"))]
49use crate::provider_capability::serial_capabilities;
50use crate::provider_capability::{engine_worker_capabilities, CpuProviderExecutionCapabilities};
51use crate::resource_domain::CpuResourceDomain;
52use crate::{CpuDomainExecutorError, CpuDomainId, CpuInnerParallelism, CpuSet};
53
54/// Operand named by a provider capability reason.
55///
56/// # Examples
57///
58/// ```
59/// use tenferro_cpu::provider::CpuOperand;
60/// assert_eq!(CpuOperand::Lhs, CpuOperand::Lhs);
61/// ```
62#[derive(Clone, Copy, Debug, PartialEq, Eq)]
63pub enum CpuOperand {
64    /// Left input operand.
65    Lhs,
66    /// Right input operand.
67    Rhs,
68    /// Writable output operand.
69    Output,
70}
71
72/// Allocation-free reason that a provider cannot execute a validated request.
73///
74/// An unsupported outcome must be reported before the provider mutates the
75/// output.
76///
77/// # Examples
78///
79/// ```
80/// use tenferro_cpu::provider::CpuProviderUnsupported;
81/// assert_eq!(
82///     CpuProviderUnsupported::RuntimeUnavailable,
83///     CpuProviderUnsupported::RuntimeUnavailable,
84/// );
85/// ```
86#[derive(Clone, Copy, Debug, PartialEq, Eq)]
87#[non_exhaustive]
88pub enum CpuProviderUnsupported {
89    /// The provider does not implement the scalar dtype.
90    DType(DType),
91    /// The provider does not implement the input ranks.
92    Rank {
93        /// Left input rank.
94        lhs: usize,
95        /// Right input rank.
96        rhs: usize,
97    },
98    /// The provider does not implement an operand layout.
99    Layout(CpuOperand),
100    /// The provider cannot implement the requested conjugation.
101    Conjugation,
102    /// The provider cannot implement the requested alpha/beta update.
103    Accumulation,
104    /// The provider cannot implement a strided batch.
105    StridedBatch,
106    /// The provider cannot implement grouped GEMM.
107    Grouped,
108    /// The optional provider runtime is not available in this process.
109    RuntimeUnavailable,
110}
111
112/// Result of attempting a validated request through one provider slot.
113///
114/// # Examples
115///
116/// ```
117/// use tenferro_cpu::provider::{CpuProviderOutcome, CpuProviderUnsupported};
118/// let outcome = CpuProviderOutcome::Unsupported(
119///     CpuProviderUnsupported::RuntimeUnavailable,
120/// );
121/// assert!(matches!(outcome, CpuProviderOutcome::Unsupported(_)));
122/// ```
123#[derive(Clone, Copy, Debug, PartialEq, Eq)]
124#[must_use]
125pub enum CpuProviderOutcome {
126    /// The provider fully executed the request.
127    Executed,
128    /// The provider did not mutate the output and resolution may continue.
129    Unsupported(CpuProviderUnsupported),
130}
131
132/// Parallel scheduling mode selected for one CPU operation.
133///
134/// # Examples
135///
136/// ```
137/// use tenferro_cpu::provider::ParallelMode;
138/// assert_ne!(ParallelMode::Sequential, ParallelMode::Outer);
139/// assert_ne!(ParallelMode::Outer, ParallelMode::Inner);
140/// ```
141#[derive(Clone, Copy, Debug, PartialEq, Eq)]
142pub enum ParallelMode {
143    /// Neither the engine nor the provider may fan out this operation.
144    Sequential,
145    /// The engine owns outer fan-out and delegates Sequential child contexts.
146    /// Providers do not receive an Outer context.
147    Outer,
148    /// One provider kernel may use the selected executor's inner region.
149    Inner,
150}
151
152/// Borrowed execution policy for an already-entered CPU operation.
153///
154/// The context exposes immutable domain facts while keeping the resource lease,
155/// executor object, and the checked executor-entry boundary private. Providers
156/// cannot install or submit work through this value.
157///
158/// # Examples
159///
160/// Providers inspect this value inside a trait method:
161///
162/// ```
163/// use tenferro_cpu::provider::CpuExecutionContext;
164/// # fn inspect(context: &CpuExecutionContext<'_>) {
165/// assert!(context.thread_budget().get() >= 1);
166/// # }
167/// ```
168#[derive(Clone, Copy)]
169pub struct CpuExecutionContext<'a> {
170    domain: &'a CpuResourceDomain,
171    parallel_mode: ParallelMode,
172    // Set only for a child of tenferro's own outer fan-out (`submit_outer` or
173    // `with_outer_lanes`), never inferred from running on a Rayon worker.
174    outer_fan_out: bool,
175    batch_policy: crate::CpuBatchPolicy,
176}
177
178impl fmt::Debug for CpuExecutionContext<'_> {
179    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
180        formatter
181            .debug_struct("CpuExecutionContext")
182            .field("domain_id", &self.domain_id())
183            .field("cpus", &self.cpus())
184            .field("thread_budget", &self.thread_budget())
185            .field("parallel_mode", &self.parallel_mode())
186            .field("outer_fan_out", &self.outer_fan_out)
187            .finish_non_exhaustive()
188    }
189}
190
191impl<'a> CpuExecutionContext<'a> {
192    fn entered(
193        domain: &'a CpuResourceDomain,
194        parallel_mode: ParallelMode,
195        batch_policy: crate::CpuBatchPolicy,
196    ) -> Self {
197        Self {
198            domain,
199            parallel_mode,
200            outer_fan_out: false,
201            batch_policy,
202        }
203    }
204
205    /// A child of tenferro's outer fan-out: sequential, with fan-out active.
206    fn outer_child(domain: &'a CpuResourceDomain, batch_policy: crate::CpuBatchPolicy) -> Self {
207        Self {
208            domain,
209            parallel_mode: ParallelMode::Sequential,
210            outer_fan_out: true,
211            batch_policy,
212        }
213    }
214
215    /// The effective batch policy for batched work in this context.
216    ///
217    /// It is the backend default unless a session scope overrides it; see
218    /// [`crate::with_batch_policy`].
219    ///
220    /// # Examples
221    ///
222    /// ```
223    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend, CpuBatchStrategy};
224    /// use tenferro_tensor::BackendSessionHost;
225    ///
226    /// let mut backend = CpuBackend::with_threads(1)?;
227    /// let strategy = backend.with_backend_session(|session| {
228    ///     with_cpu_exec_session(session, |cpu| {
229    ///         cpu.with_linalg_pool(|context, _| Ok(context.batch_policy().strategy()))
230    ///     })
231    ///     .expect("a CPU backend session")
232    /// })??;
233    /// assert_eq!(strategy, CpuBatchStrategy::Auto);
234    /// # Ok::<(), Box<dyn std::error::Error>>(())
235    /// ```
236    #[must_use]
237    pub fn batch_policy(&self) -> crate::CpuBatchPolicy {
238        self.batch_policy
239    }
240
241    /// This context with a scoped batch-policy override.
242    pub(crate) fn with_batch_policy(mut self, batch_policy: crate::CpuBatchPolicy) -> Self {
243        self.batch_policy = batch_policy;
244        self
245    }
246
247    /// Whether [`Self::with_outer_lanes`] can fan out from this context.
248    ///
249    /// # Examples
250    ///
251    /// ```
252    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
253    /// use tenferro_tensor::BackendSessionHost;
254    ///
255    /// let mut backend = CpuBackend::with_threads(1)?;
256    /// let fans_out = backend.with_backend_session(|session| {
257    ///     with_cpu_exec_session(session, |cpu| {
258    ///         cpu.with_linalg_pool(|context, _| Ok(context.can_fan_out_lanes()))
259    ///     })
260    ///     .expect("a CPU backend session")
261    /// })??;
262    /// assert!(!fans_out);
263    /// # Ok::<(), Box<dyn std::error::Error>>(())
264    /// ```
265    #[must_use]
266    pub fn can_fan_out_lanes(&self) -> bool {
267        self.parallel_mode == ParallelMode::Inner
268            && self.domain.executor_capabilities().inner_parallelism == CpuInnerParallelism::Rayon
269            && self.thread_budget().get() > 1
270    }
271
272    /// Whether this context is one lane of tenferro's own outer fan-out.
273    ///
274    /// Such a lane runs concurrently with its siblings. Its mode is
275    /// [`ParallelMode::Sequential`], so it may not start inner parallel work, and
276    /// a provider whose declared parallelism is an independent runtime (the
277    /// default for external BLAS/LAPACK) is rejected before dispatch. Running on
278    /// a Rayon worker is not by itself outer fan-out.
279    ///
280    /// # Examples
281    ///
282    /// ```
283    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
284    /// use tenferro_tensor::BackendSessionHost;
285    ///
286    /// let mut backend = CpuBackend::with_threads(1)?;
287    /// let lane = backend.with_backend_session(|session| {
288    ///     with_cpu_exec_session(session, |cpu| {
289    ///         cpu.with_linalg_pool(|context, _| Ok(context.is_outer_fan_out_lane()))
290    ///     })
291    ///     .expect("a CPU backend session")
292    /// })??;
293    /// assert!(!lane);
294    /// # Ok::<(), Box<dyn std::error::Error>>(())
295    /// ```
296    #[must_use]
297    pub fn is_outer_fan_out_lane(&self) -> bool {
298        self.outer_fan_out
299    }
300
301    /// Run independent lane jobs as tenferro-owned outer fan-out inside this
302    /// already-entered context's own Rayon region.
303    ///
304    /// Each item of `jobs` (typically one disjoint chunk of a batch) runs once
305    /// and receives a lane context: [`Self::is_outer_fan_out_lane`] is true and
306    /// the mode is [`ParallelMode::Sequential`], so the lane neither starts
307    /// inner parallel work nor reaches an independent-runtime provider. Fan-out
308    /// needs an [`ParallelMode::Inner`] context of a Rayon executor with more
309    /// than one thread; otherwise the jobs run in order on the calling thread,
310    /// still with a lane context. No second pool is created: jobs run on the
311    /// Rayon pool this context was entered in.
312    ///
313    /// # Examples
314    ///
315    /// ```
316    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
317    /// use tenferro_tensor::BackendSessionHost;
318    ///
319    /// let mut backend = CpuBackend::with_threads(2)?;
320    /// let mut data = vec![1.0_f64; 8];
321    /// backend.with_backend_session(|session| {
322    ///     with_cpu_exec_session(session, |cpu| {
323    ///         cpu.with_linalg_pool(|context, _| {
324    ///             context.with_outer_lanes(data.chunks_mut(3), |chunk, lane| {
325    ///                 assert!(lane.is_outer_fan_out_lane());
326    ///                 chunk.iter_mut().for_each(|value| *value *= 2.0);
327    ///             });
328    ///             Ok(())
329    ///         })
330    ///     })
331    ///     .expect("a CPU backend session")
332    /// })??;
333    /// assert_eq!(data, [2.0; 8]);
334    /// # Ok::<(), Box<dyn std::error::Error>>(())
335    /// ```
336    pub fn with_outer_lanes<I>(
337        &self,
338        jobs: I,
339        job: impl Fn(I::Item, &CpuExecutionContext<'a>) + Sync,
340    ) where
341        I: IntoIterator,
342        I::IntoIter: Send,
343        I::Item: Send,
344    {
345        let child = Self::outer_child(self.domain, self.batch_policy);
346        if !self.can_fan_out_lanes() {
347            for item in jobs {
348                job(item, &child);
349            }
350            return;
351        }
352        let job = &job;
353        let jobs = jobs.into_iter();
354        rayon::scope(|scope| {
355            for item in jobs {
356                scope.spawn(move |_| job(item, &child));
357            }
358        });
359    }
360
361    /// Return the stable identity of the selected CPU resource domain.
362    ///
363    /// # Examples
364    ///
365    /// ```
366    /// use tenferro_cpu::CpuExecutionContext;
367    /// # fn inspect(context: &CpuExecutionContext<'_>) {
368    /// let _domain_id = context.domain_id();
369    /// # }
370    /// ```
371    pub fn domain_id(&self) -> CpuDomainId {
372        self.domain.id()
373    }
374
375    /// Return the selected domain's declared logical CPU set, when present.
376    ///
377    /// # Examples
378    ///
379    /// ```
380    /// use tenferro_cpu::CpuExecutionContext;
381    /// # fn inspect(context: &CpuExecutionContext<'_>) {
382    /// if let Some(cpus) = context.cpus() {
383    ///     assert!(!cpus.is_empty());
384    /// }
385    /// # }
386    /// ```
387    pub fn cpus(&self) -> Option<&CpuSet> {
388        self.domain.cpus()
389    }
390
391    /// Return the selected domain's admission contract.
392    ///
393    /// # Examples
394    ///
395    /// ```rust
396    /// use tenferro_cpu::{CpuAdmissionMode, CpuExecutionContext};
397    /// # fn inspect(context: &CpuExecutionContext<'_>) {
398    /// let _mode: CpuAdmissionMode = context.admission_mode();
399    /// # }
400    /// ```
401    pub fn admission_mode(&self) -> crate::CpuAdmissionMode {
402        self.domain.admission_mode()
403    }
404
405    /// Return the non-zero maximum participating-thread budget.
406    ///
407    /// # Examples
408    ///
409    /// ```
410    /// use tenferro_cpu::CpuExecutionContext;
411    /// # fn inspect(context: &CpuExecutionContext<'_>) {
412    /// assert!(context.thread_budget().get() >= 1);
413    /// # }
414    /// ```
415    pub fn thread_budget(&self) -> NonZeroUsize {
416        self.domain.thread_budget()
417    }
418
419    /// Return the engine-selected scheduling mode for this entered provider call.
420    ///
421    /// Provider calls observe Sequential or Inner. Outer scheduling creates a
422    /// separate Sequential context inside every submitted child.
423    ///
424    /// # Examples
425    ///
426    /// ```
427    /// use tenferro_cpu::{CpuExecutionContext, ParallelMode};
428    /// # fn inspect(context: &CpuExecutionContext<'_>) {
429    /// assert!(matches!(
430    ///     context.parallel_mode(),
431    ///     ParallelMode::Sequential | ParallelMode::Inner
432    /// ));
433    /// # }
434    /// ```
435    pub fn parallel_mode(&self) -> ParallelMode {
436        self.parallel_mode
437    }
438
439    /// Materialize a borrowed tensor view for one scoped operation and reclaim
440    /// its temporary host buffer before returning.
441    ///
442    /// Owned tensor inputs are borrowed directly. View inputs are materialized
443    /// from `buffers`, passed to `operation`, and returned to the same pool on
444    /// both success and ordinary error. The receiver is an unforgeable proof
445    /// that the caller is already inside the selected CPU execution domain.
446    ///
447    /// # Errors
448    ///
449    /// Returns [`tenferro_tensor::Error::RuntimeState`] when the view is not
450    /// accessible from CPU host memory, propagates typed view-materialization
451    /// errors, and otherwise returns the error produced by `operation`.
452    ///
453    /// # Examples
454    ///
455    /// ```
456    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
457    /// use tenferro_tensor::{
458    ///     BackendSessionHost, StridedSliceSpec, TensorRead, TensorView, TypedTensor,
459    /// };
460    ///
461    /// let mut backend = CpuBackend::with_threads(1)?;
462    /// let input = TypedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
463    /// backend.with_backend_session(|session| {
464    ///     with_cpu_exec_session(session, |cpu| {
465    ///         cpu.with_linalg_pool(|context, buffers| {
466    ///             let view = input
467    ///                 .as_view()
468    ///                 .slice_view(&[StridedSliceSpec::reverse()])?;
469    ///             context.with_materialized_tensor_read(
470    ///                 buffers,
471    ///                 "example",
472    ///                 TensorRead::from_view(TensorView::F64(view)),
473    ///                 |materialized, _| {
474    ///                     assert_eq!(materialized.as_slice::<f64>().unwrap(), &[2.0, 1.0]);
475    ///                     Ok(())
476    ///                 },
477    ///             )
478    ///         })
479    ///     })
480    ///     .expect("a CPU backend session")
481    /// })??;
482    /// # Ok::<(), Box<dyn std::error::Error>>(())
483    /// ```
484    #[doc(hidden)]
485    pub fn with_materialized_tensor_read<R>(
486        &self,
487        buffers: &mut BufferPool,
488        op: &'static str,
489        input: TensorRead<'_>,
490        operation: impl FnOnce(&Tensor, &mut BufferPool) -> tenferro_tensor::Result<R>,
491    ) -> tenferro_tensor::Result<R> {
492        match input {
493            TensorRead::Tensor(tensor) => operation(tensor, buffers),
494            TensorRead::View(view) => {
495                let materialized = self.with_native_parallelism(|| {
496                    crate::materialize_tensor_read(buffers, op, TensorRead::View(view))
497                })?;
498                let result = operation(&materialized, buffers);
499                crate::backend::reclaim_tensor(buffers, materialized);
500                result
501            }
502        }
503    }
504
505    /// Reshape a compact tensor while retaining the current execution proof.
506    ///
507    /// This metadata-only helper does not enter an executor or borrow another
508    /// scratch pool.
509    ///
510    /// # Errors
511    ///
512    /// Returns [`tenferro_tensor::Error::Validation`] when the input and output
513    /// shapes have different element counts or the requested layout is invalid.
514    ///
515    /// # Examples
516    ///
517    /// ```
518    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
519    /// use tenferro_tensor::{BackendSessionHost, Tensor};
520    ///
521    /// let mut backend = CpuBackend::with_threads(1)?;
522    /// let input = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
523    /// backend.with_backend_session(|session| {
524    ///     with_cpu_exec_session(session, |cpu| {
525    ///         cpu.with_linalg_pool(|context, _| {
526    ///             let output = context.reshape_tensor(&input, &[2, 1])?;
527    ///             assert_eq!(output.shape(), &[2, 1]);
528    ///             Ok(())
529    ///         })
530    ///     })
531    ///     .expect("a CPU backend session")
532    /// })??;
533    /// # Ok::<(), Box<dyn std::error::Error>>(())
534    /// ```
535    #[doc(hidden)]
536    pub fn reshape_tensor(
537        &self,
538        input: &Tensor,
539        shape: &[usize],
540    ) -> tenferro_tensor::Result<Tensor> {
541        crate::structural::reshape(input, shape)
542    }
543
544    /// Return the faer policy selected by this operation context.
545    ///
546    /// This hidden public method is the owner-scoped extension contract used by
547    /// operation-family crates such as `tenferro-linalg`. Keeping the mapping
548    /// here prevents sibling crates from deriving a second CPU threading
549    /// policy.
550    ///
551    /// # Examples
552    ///
553    /// ```
554    /// use tenferro_cpu::CpuExecutionContext;
555    /// # fn inspect(context: &CpuExecutionContext<'_>) {
556    /// let _policy = context.faer_parallelism();
557    /// # }
558    /// ```
559    #[cfg(feature = "cpu-faer")]
560    #[doc(hidden)]
561    pub fn faer_parallelism(self) -> faer::Par {
562        match (
563            self.parallel_mode,
564            self.domain.executor_capabilities().inner_parallelism,
565        ) {
566            (ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
567                faer::Par::rayon(self.thread_budget().get())
568            }
569            _ => faer::Par::Seq,
570        }
571    }
572
573    /// The Rayon pool this context's inner parallel region runs on.
574    ///
575    /// `Some` exactly when the context owns an inner region: parallel mode
576    /// [`ParallelMode::Inner`], a Rayon-backed executor, and a thread budget
577    /// above one (the same gate as the faer policy). The caller is already on
578    /// a worker of this pool, so work installed or scoped on it runs in
579    /// place. A provider running its own kernels on the pool must use at most
580    /// [`CpuExecutionContext::thread_budget`] threads, which can be smaller
581    /// than the pool, and declares
582    /// [`crate::CpuThreadCountControl::PerCallUpperBound`] with
583    /// [`crate::CpuPlacementControl::EngineWorkers`].
584    ///
585    /// # Examples
586    ///
587    /// ```
588    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
589    /// use tenferro_tensor::BackendSessionHost;
590    /// let mut backend = CpuBackend::with_threads(2)?;
591    /// let workers = backend.with_backend_session(|session| {
592    ///     with_cpu_exec_session(session, |cpu| {
593    ///         cpu.with_linalg_pool(|context, _| {
594    ///             Ok(context.rayon_pool().map(|pool| pool.current_num_threads()))
595    ///         })
596    ///     })
597    ///     .expect("a CPU backend session")
598    /// })??;
599    /// assert_eq!(workers, Some(2));
600    /// # Ok::<(), Box<dyn std::error::Error>>(())
601    /// ```
602    pub fn rayon_pool(&self) -> Option<&'a rayon::ThreadPool> {
603        match (
604            self.parallel_mode,
605            self.domain.executor_capabilities().inner_parallelism,
606        ) {
607            (ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
608                self.domain.executor().rayon_pool()
609            }
610            _ => None,
611        }
612    }
613
614    /// Effective native-kernel degree inside this already-entered CPU context.
615    /// Non-Rayon executors and sequential/nested policy use one thread.
616    ///
617    /// # Examples
618    /// ```
619    /// use tenferro_cpu::{with_cpu_exec_session, CpuBackend};
620    /// use tenferro_tensor::BackendSessionHost;
621    /// let mut backend = CpuBackend::with_threads(1)?;
622    /// let threads = backend.with_backend_session(|session| {
623    ///     with_cpu_exec_session(session, |cpu| {
624    ///         cpu.with_linalg_pool(|context, _| Ok(context.native_thread_count()))
625    ///     })
626    ///     .expect("a CPU backend session")
627    /// })??;
628    /// assert_eq!(threads, 1);
629    /// # Ok::<(), Box<dyn std::error::Error>>(())
630    /// ```
631    #[doc(hidden)]
632    pub fn native_thread_count(&self) -> usize {
633        match (
634            self.parallel_mode,
635            self.domain.executor_capabilities().inner_parallelism,
636        ) {
637            (ParallelMode::Inner, CpuInnerParallelism::Rayon) => self.thread_budget().get(),
638            _ => 1,
639        }
640    }
641
642    pub(crate) fn strided_exec_context(&self) -> strided_kernel::ExecContext {
643        match (
644            self.parallel_mode,
645            self.domain.executor_capabilities().inner_parallelism,
646        ) {
647            (ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
648                // This is only an operation-local thread limit for strided's
649                // replay policy. The Rayon pool itself is the already-entered
650                // CpuContext pool installed by `with_native_parallelism`.
651                match strided_kernel::ExecContext::max_threads(self.thread_budget().get()) {
652                    Ok(context) => context,
653                    // INVARIANT: CpuExecutionContext stores a NonZeroUsize
654                    // thread budget, and this branch passes that positive value.
655                    Err(_) => unreachable!("CpuExecutionContext has a non-zero thread budget"),
656                }
657            }
658            _ => strided_kernel::ExecContext::serial(),
659        }
660    }
661
662    pub(crate) fn with_native_parallelism<R>(&self, operation: impl FnOnce() -> R) -> R {
663        let policy = match (
664            self.parallel_mode,
665            self.domain.executor_capabilities().inner_parallelism,
666        ) {
667            (ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
668                strided_kernel::ExecutionPolicy::Rayon {
669                    max_threads: self.thread_budget(),
670                }
671            }
672            _ => strided_kernel::ExecutionPolicy::Sequential,
673        };
674        strided_kernel::with_execution_policy(policy, operation)
675    }
676}
677
678/// Proof that every provider an outer fan-out's lanes can reach accepts
679/// concurrent sequential calls.
680///
681/// [`CpuOperationEntry::submit_outer`] takes this value, so fan-out cannot be
682/// submitted before the reachable delegates' declarations are checked: a
683/// declared nesting violation fails before any lane runs or writes output.
684#[derive(Debug)]
685pub(crate) struct OuterFanOutChecked(());
686
687/// Check the reachable delegates of an outer fan-out before submitting it.
688///
689/// Every capability listed must accept [`ParallelMode::Outer`] (sequential,
690/// concurrent-call safe, worker-local). An implementation whose parallelism is
691/// an independent runtime, the default declaration for external BLAS/LAPACK,
692/// is rejected.
693pub(crate) fn check_outer_fan_out_delegates<'c>(
694    delegates: impl IntoIterator<Item = &'c crate::CpuProviderExecutionCapabilities>,
695) -> Result<OuterFanOutChecked, crate::CpuProviderDomainError> {
696    for capabilities in delegates {
697        if !capabilities.accepts_mode(ParallelMode::Outer) {
698            return Err(crate::CpuProviderDomainError::ParallelModeNotSupported {
699                mode: ParallelMode::Outer,
700            });
701        }
702    }
703    Ok(OuterFanOutChecked(()))
704}
705
706/// Where an outer fan-out runs: submitted to the domain executor from an
707/// unentered operation, or as lanes of an already-entered Inner context.
708#[derive(Clone, Copy)]
709pub(crate) enum CpuOuterFanOut<'a> {
710    Executor(CpuOperationEntry<'a>),
711    Lanes(CpuExecutionContext<'a>),
712}
713
714impl CpuOuterFanOut<'_> {
715    /// Run `len` indexed jobs, each with a lane context, and return the first
716    /// job error.
717    pub(crate) fn submit(
718        self,
719        checked: OuterFanOutChecked,
720        len: usize,
721        operation: impl Fn(usize, &CpuExecutionContext<'_>) -> Result<(), CpuDomainExecutorError> + Sync,
722    ) -> Result<(), CpuDomainExecutorError> {
723        match self {
724            Self::Executor(entry) => entry.submit_outer(checked, len, operation),
725            Self::Lanes(context) => {
726                let first_error = std::sync::Mutex::new(None);
727                context.with_outer_lanes(0..len, |index, lane| {
728                    if let Err(error) = operation(index, lane) {
729                        first_error
730                            .lock()
731                            .unwrap_or_else(std::sync::PoisonError::into_inner)
732                            .get_or_insert(error);
733                    }
734                });
735                match first_error
736                    .into_inner()
737                    .unwrap_or_else(std::sync::PoisonError::into_inner)
738                {
739                    Some(error) => Err(error),
740                    None => Ok(()),
741                }
742            }
743        }
744    }
745}
746
747/// Crate-private unentered capability for one CPU operation.
748///
749/// This is the only type that owns the resource permit and may cross the
750/// selected domain executor boundary. A [`CpuExecutionContext`] is constructed
751/// only inside an installed job or an outer child job.
752#[derive(Clone, Copy)]
753pub(crate) struct CpuOperationEntry<'a> {
754    domain: &'a CpuResourceDomain,
755    permit: &'a ResourcePermit,
756    batch_policy: crate::CpuBatchPolicy,
757}
758
759impl<'a> CpuOperationEntry<'a> {
760    pub(crate) fn new(domain: &'a CpuResourceDomain, permit: &'a ResourcePermit) -> Self {
761        Self {
762            domain,
763            permit,
764            batch_policy: crate::CpuBatchPolicy::default(),
765        }
766    }
767
768    /// This entry with a different effective batch policy.
769    pub(crate) fn with_batch_policy(mut self, batch_policy: crate::CpuBatchPolicy) -> Self {
770        self.batch_policy = batch_policy;
771        self
772    }
773
774    pub(crate) fn batch_policy(self) -> crate::CpuBatchPolicy {
775        self.batch_policy
776    }
777
778    pub(crate) fn domain_id(self) -> CpuDomainId {
779        self.domain.id()
780    }
781
782    pub(crate) fn enter<R: Send>(
783        self,
784        parallel_mode: ParallelMode,
785        operation: impl FnOnce(&CpuExecutionContext<'_>) -> R + Send,
786    ) -> Result<R, CpuDomainExecutorError> {
787        if parallel_mode == ParallelMode::Outer {
788            return Err(CpuDomainExecutorError::Scheduling {
789                message: "CPU executor install requires Sequential or Inner mode, got Outer"
790                    .to_owned(),
791            });
792        }
793        let owner = self.permit.owner();
794        if crate::backend::execution_scope::is_entered(self.domain, self.permit) {
795            return Ok(with_execution_owner(owner, || {
796                let context =
797                    CpuExecutionContext::entered(self.domain, parallel_mode, self.batch_policy);
798                operation(&context)
799            }));
800        }
801        with_execution_owner(owner, || {
802            install_scoped(self.domain.executor().as_ref(), || {
803                with_execution_owner(owner, || {
804                    let context =
805                        CpuExecutionContext::entered(self.domain, parallel_mode, self.batch_policy);
806                    operation(&context)
807                })
808            })
809        })
810    }
811
812    pub(crate) fn enter_or_reuse<R: Send>(
813        self,
814        entered: Option<&CpuExecutionContext<'_>>,
815        parallel_mode: ParallelMode,
816        operation: impl FnOnce(&CpuExecutionContext<'_>) -> R + Send,
817    ) -> Result<R, CpuDomainExecutorError> {
818        let Some(entered) = entered else {
819            return self.enter(parallel_mode, operation);
820        };
821        if parallel_mode == ParallelMode::Outer {
822            return Err(CpuDomainExecutorError::Scheduling {
823                message: "entered CPU session requires Sequential or Inner mode, got Outer"
824                    .to_owned(),
825            });
826        }
827        if entered.domain_id() != self.domain.id() {
828            return Err(CpuDomainExecutorError::Scheduling {
829                message: format!(
830                    "entered CPU session domain {:?} does not match operation domain {:?}",
831                    entered.domain_id(),
832                    self.domain.id()
833                ),
834            });
835        }
836        let owner = self.permit.owner();
837        // A lane of outer fan-out keeps its fan-out fact, and it may not widen
838        // its sequential policy into inner parallelism: that would nest a second
839        // fan-out inside every sibling lane.
840        let context = if entered.is_outer_fan_out_lane() {
841            CpuExecutionContext::outer_child(self.domain, self.batch_policy)
842        } else {
843            CpuExecutionContext::entered(self.domain, parallel_mode, self.batch_policy)
844        };
845        Ok(with_execution_owner(owner, || operation(&context)))
846    }
847
848    /// Whether a backend session enters this domain's executor once for its
849    /// whole callback. Externally managed domains enter per operation instead.
850    pub(crate) fn enters_executor_per_session(self) -> bool {
851        self.domain.ownership() == crate::CpuDomainOwnership::Managed
852    }
853
854    /// Enter a Tenferro-managed executor once for a whole backend session.
855    ///
856    /// Callers check [`Self::enters_executor_per_session`] first; the managed
857    /// Rayon executor's synchronous install does not fail for Sequential or
858    /// Inner mode, but its typed error is still reported rather than hidden.
859    pub(crate) fn enter_managed_session<R: Send>(
860        self,
861        operation: impl FnOnce(CpuExecutionContext<'a>) -> R + Send,
862    ) -> Result<R, tenferro_tensor::SessionEntryError> {
863        let mode = self.preferred_engine_mode();
864        self.enter(mode, |_| {
865            operation(CpuExecutionContext::entered(
866                self.domain,
867                mode,
868                self.batch_policy,
869            ))
870        })
871        .map_err(|error| tenferro_tensor::SessionEntryError::Executor {
872            backend: crate::backend::CPU_BACKEND,
873            source: Box::new(error),
874        })
875    }
876
877    pub(crate) fn submit_outer(
878        self,
879        _delegates: OuterFanOutChecked,
880        len: usize,
881        operation: impl Fn(usize, &CpuExecutionContext<'_>) -> Result<(), CpuDomainExecutorError> + Sync,
882    ) -> Result<(), CpuDomainExecutorError> {
883        if !self.supports_outer() {
884            return Err(CpuDomainExecutorError::Scheduling {
885                message: format!(
886                    "CPU domain {:?} does not support Outer mode",
887                    self.domain.id()
888                ),
889            });
890        }
891        let owner = self.permit.owner();
892        let lane_count = len.min(self.domain.thread_budget().get());
893        let jobs = indexed_jobs(lane_count, |lane| {
894            // INVARIANT: valid lanes partition `0..len` by residue modulo the
895            // nonzero `lane_count`, so every logical job runs exactly once
896            // while the executor can schedule at most the domain budget.
897            let mut index = lane;
898            while index < len {
899                with_execution_owner(owner, || {
900                    let context = CpuExecutionContext::outer_child(self.domain, self.batch_policy);
901                    operation(index, &context)
902                })?;
903                let Some(next) = index.checked_add(lane_count) else {
904                    break;
905                };
906                index = next;
907            }
908            Ok(())
909        });
910        with_execution_owner(owner, || self.domain.executor().submit(&jobs))?;
911        if let Some(index) = jobs.invalid_index_attempt() {
912            return Err(CpuDomainExecutorError::Scheduling {
913                message: format!(
914                    "executor requested scoped CPU lane index {index}, but the submission has {lane_count} lanes for {len} logical jobs"
915                ),
916            });
917        }
918        Ok(())
919    }
920
921    pub(crate) fn preferred_engine_mode(self) -> ParallelMode {
922        if self.domain.thread_budget().get() > 1
923            && self.domain.executor_capabilities().inner_parallelism == CpuInnerParallelism::Rayon
924        {
925            ParallelMode::Inner
926        } else {
927            ParallelMode::Sequential
928        }
929    }
930
931    pub(crate) fn preferred_provider_mode(
932        self,
933        accepts: impl Fn(ParallelMode) -> bool,
934    ) -> Result<ParallelMode, crate::CpuProviderDomainError> {
935        if self.domain.thread_budget().get() == 1 {
936            return if accepts(ParallelMode::Sequential) {
937                Ok(ParallelMode::Sequential)
938            } else {
939                Err(crate::CpuProviderDomainError::ParallelModeNotSupported {
940                    mode: ParallelMode::Sequential,
941                })
942            };
943        }
944        if accepts(ParallelMode::Inner) {
945            return Ok(ParallelMode::Inner);
946        }
947        if accepts(ParallelMode::Sequential) {
948            return Ok(ParallelMode::Sequential);
949        }
950        Err(crate::CpuProviderDomainError::ParallelModeNotSupported {
951            mode: ParallelMode::Inner,
952        })
953    }
954
955    pub(crate) fn provider_default_compatibility_mode(self) -> ParallelMode {
956        if self.domain.thread_budget().get() == 1 {
957            ParallelMode::Sequential
958        } else {
959            ParallelMode::Inner
960        }
961    }
962
963    pub(crate) fn thread_budget(self) -> NonZeroUsize {
964        self.domain.thread_budget()
965    }
966
967    pub(crate) fn supports_outer(self) -> bool {
968        self.domain.thread_budget().get() > 1
969            && self.domain.executor_capabilities().outer_parallelism
970    }
971}
972
973/// Checked element offset and strides for a batched matrix operand.
974///
975/// # Examples
976///
977/// Providers receive this descriptor from a validated request:
978///
979/// ```
980/// use tenferro_cpu::provider::CpuGemmRequest;
981/// # fn inspect(request: &CpuGemmRequest<'_, '_, '_>) {
982/// assert!(request.lhs_layout().row_stride() != 0 || request.rows() <= 1);
983/// # }
984/// ```
985#[derive(Clone, Copy, Debug, PartialEq, Eq)]
986pub struct CpuBatchedMatrixLayout {
987    offset: isize,
988    row_stride: isize,
989    column_stride: isize,
990    batch_stride: isize,
991}
992
993impl CpuBatchedMatrixLayout {
994    #[allow(dead_code)]
995    pub(crate) fn new(
996        offset: isize,
997        row_stride: isize,
998        column_stride: isize,
999        batch_stride: isize,
1000    ) -> Self {
1001        Self {
1002            offset,
1003            row_stride,
1004            column_stride,
1005            batch_stride,
1006        }
1007    }
1008
1009    /// Return the checked base element offset.
1010    pub fn offset(self) -> isize {
1011        self.offset
1012    }
1013
1014    /// Return the row element stride.
1015    pub fn row_stride(self) -> isize {
1016        self.row_stride
1017    }
1018
1019    /// Return the column element stride.
1020    pub fn column_stride(self) -> isize {
1021        self.column_stride
1022    }
1023
1024    /// Return the batch element stride.
1025    pub fn batch_stride(self) -> isize {
1026        self.batch_stride
1027    }
1028}
1029
1030/// Whether a provider may, must, or must not execute a batch of GEMMs with
1031/// one vendor batch call (`cblas_?gemm_batch`).
1032///
1033/// The engine resolves this from the effective [`crate::CpuBatchPolicy`]
1034/// before any output write. A provider without a vendor batch routine treats
1035/// [`Self::Allowed`] and [`Self::Forbidden`] alike and reports
1036/// [`CpuProviderUnsupported::RuntimeUnavailable`] for [`Self::Required`].
1037///
1038/// # Examples
1039///
1040/// ```
1041/// use tenferro_cpu::provider::CpuVendorBatch;
1042///
1043/// let allowed = CpuVendorBatch::Allowed { max_item_dim: 16 };
1044/// assert!(allowed.permits([[4, 4, 4], [8, 8, 8]]));
1045/// assert!(!allowed.permits([[4, 4, 32], [8, 8, 8]]));
1046/// assert!(!CpuVendorBatch::Forbidden.permits([[1, 1, 1], [1, 1, 1]]));
1047/// ```
1048#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1049#[non_exhaustive]
1050pub enum CpuVendorBatch {
1051    /// Use a vendor batch call when there is more than one item and every
1052    /// item's `m`, `n` and `k` are at most `max_item_dim`.
1053    Allowed {
1054        /// Largest per-item dimension for which a vendor batch call is used.
1055        max_item_dim: usize,
1056    },
1057    /// Execute the whole batch with one vendor batch call.
1058    Required,
1059    /// Never use a vendor batch call.
1060    Forbidden,
1061}
1062
1063impl Default for CpuVendorBatch {
1064    fn default() -> Self {
1065        Self::Allowed {
1066            max_item_dim: crate::CpuBatchThresholds::default().vendor_batch_max_item_dim(),
1067        }
1068    }
1069}
1070
1071impl CpuVendorBatch {
1072    /// Whether this control selects a vendor batch call for GEMMs of the given
1073    /// `[m, n, k]` dimensions; [`Self::Required`] always does.
1074    ///
1075    /// # Examples
1076    ///
1077    /// ```
1078    /// use tenferro_cpu::provider::CpuVendorBatch;
1079    /// assert!(CpuVendorBatch::Required.permits([[64, 64, 64]]));
1080    /// ```
1081    #[must_use]
1082    pub fn permits(self, dims: impl IntoIterator<Item = [usize; 3]>) -> bool {
1083        match self {
1084            Self::Allowed { max_item_dim } => crate::CpuBatchThresholds::default()
1085                .with_vendor_batch_max_item_dim(max_item_dim)
1086                .auto_uses_vendor_batch(dims),
1087            Self::Required => true,
1088            Self::Forbidden => false,
1089        }
1090    }
1091}
1092
1093/// Validated borrowed GEMM request.
1094///
1095/// A batch count of one is a single GEMM. A larger batch count is a strided
1096/// batched GEMM.
1097///
1098/// # Examples
1099///
1100/// ```
1101/// use tenferro_cpu::provider::CpuGemmRequest;
1102/// # fn inspect(request: &CpuGemmRequest<'_, '_, '_>) {
1103/// assert!(request.batch_count() >= 1);
1104/// # }
1105/// ```
1106#[derive(Debug)]
1107pub struct CpuGemmRequest<'request, 'input, 'output> {
1108    lhs: &'request TensorRead<'input>,
1109    rhs: &'request TensorRead<'input>,
1110    output: &'request mut TensorWrite<'output>,
1111    rows: usize,
1112    columns: usize,
1113    contracted: usize,
1114    batch_count: usize,
1115    lhs_layout: CpuBatchedMatrixLayout,
1116    rhs_layout: CpuBatchedMatrixLayout,
1117    output_layout: CpuBatchedMatrixLayout,
1118    accumulation: DotGeneralAccumulation,
1119    vendor_batch: CpuVendorBatch,
1120}
1121
1122pub(crate) struct CpuGemmRequestParts<'request, 'input, 'output> {
1123    pub(crate) lhs: &'request TensorRead<'input>,
1124    pub(crate) rhs: &'request TensorRead<'input>,
1125    pub(crate) output: &'request mut TensorWrite<'output>,
1126    pub(crate) rows: usize,
1127    pub(crate) columns: usize,
1128    pub(crate) contracted: usize,
1129    pub(crate) batch_count: usize,
1130    pub(crate) lhs_layout: CpuBatchedMatrixLayout,
1131    pub(crate) rhs_layout: CpuBatchedMatrixLayout,
1132    pub(crate) output_layout: CpuBatchedMatrixLayout,
1133    pub(crate) accumulation: DotGeneralAccumulation,
1134    // Read only by the BLAS provider; faer checks the request accessor.
1135    #[cfg(feature = "cpu-blas")]
1136    pub(crate) vendor_batch: CpuVendorBatch,
1137}
1138
1139impl<'request, 'input, 'output> CpuGemmRequest<'request, 'input, 'output> {
1140    #[allow(clippy::too_many_arguments, dead_code)]
1141    pub(crate) fn new(
1142        lhs: &'request TensorRead<'input>,
1143        rhs: &'request TensorRead<'input>,
1144        output: &'request mut TensorWrite<'output>,
1145        rows: usize,
1146        columns: usize,
1147        contracted: usize,
1148        batch_count: usize,
1149        lhs_layout: CpuBatchedMatrixLayout,
1150        rhs_layout: CpuBatchedMatrixLayout,
1151        output_layout: CpuBatchedMatrixLayout,
1152        accumulation: DotGeneralAccumulation,
1153    ) -> Self {
1154        Self {
1155            lhs,
1156            rhs,
1157            output,
1158            rows,
1159            columns,
1160            contracted,
1161            batch_count,
1162            lhs_layout,
1163            rhs_layout,
1164            output_layout,
1165            accumulation,
1166            vendor_batch: CpuVendorBatch::default(),
1167        }
1168    }
1169
1170    /// This request with the engine-resolved vendor-batch control.
1171    pub(crate) fn with_vendor_batch(mut self, vendor_batch: CpuVendorBatch) -> Self {
1172        self.vendor_batch = vendor_batch;
1173        self
1174    }
1175
1176    /// Return whether a batch may, must or must not use a vendor batch call.
1177    ///
1178    /// # Examples
1179    ///
1180    /// ```
1181    /// use tenferro_cpu::provider::{CpuGemmRequest, CpuVendorBatch};
1182    /// # fn inspect(request: &CpuGemmRequest<'_, '_, '_>) {
1183    /// let _ = request.vendor_batch() == CpuVendorBatch::Forbidden;
1184    /// # }
1185    /// ```
1186    pub fn vendor_batch(&self) -> CpuVendorBatch {
1187        self.vendor_batch
1188    }
1189
1190    /// Return the borrowed left input.
1191    ///
1192    /// # Examples
1193    ///
1194    /// ```
1195    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1196    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1197    /// #     let _ = request.lhs();
1198    /// # }
1199    /// ```
1200    pub fn lhs(&self) -> &TensorRead<'input> {
1201        self.lhs
1202    }
1203
1204    /// Return the borrowed right input.
1205    ///
1206    /// # Examples
1207    ///
1208    /// ```
1209    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1210    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1211    /// #     let _ = request.rhs();
1212    /// # }
1213    /// ```
1214    pub fn rhs(&self) -> &TensorRead<'input> {
1215        self.rhs
1216    }
1217
1218    /// Reborrow the writable output for the duration of the current call.
1219    pub fn output(&mut self) -> &mut TensorWrite<'output> {
1220        self.output
1221    }
1222
1223    /// Return the number of output rows.
1224    ///
1225    /// # Examples
1226    ///
1227    /// ```
1228    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1229    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1230    /// #     let _ = request.rows();
1231    /// # }
1232    /// ```
1233    pub fn rows(&self) -> usize {
1234        self.rows
1235    }
1236
1237    /// Return the number of output columns.
1238    ///
1239    /// # Examples
1240    ///
1241    /// ```
1242    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1243    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1244    /// #     let _ = request.columns();
1245    /// # }
1246    /// ```
1247    pub fn columns(&self) -> usize {
1248        self.columns
1249    }
1250
1251    /// Return the contracted dimension.
1252    pub fn contracted(&self) -> usize {
1253        self.contracted
1254    }
1255
1256    /// Return the number of matrices in this strided batch.
1257    pub fn batch_count(&self) -> usize {
1258        self.batch_count
1259    }
1260
1261    /// Return the left input matrix layout.
1262    pub fn lhs_layout(&self) -> CpuBatchedMatrixLayout {
1263        self.lhs_layout
1264    }
1265
1266    /// Return the right input matrix layout.
1267    pub fn rhs_layout(&self) -> CpuBatchedMatrixLayout {
1268        self.rhs_layout
1269    }
1270
1271    /// Return the output matrix layout.
1272    pub fn output_layout(&self) -> CpuBatchedMatrixLayout {
1273        self.output_layout
1274    }
1275
1276    /// Return conjugation and alpha/beta update semantics.
1277    pub fn accumulation(&self) -> DotGeneralAccumulation {
1278        self.accumulation
1279    }
1280
1281    pub(crate) fn into_parts(self) -> CpuGemmRequestParts<'request, 'input, 'output> {
1282        CpuGemmRequestParts {
1283            #[cfg(feature = "cpu-blas")]
1284            vendor_batch: self.vendor_batch,
1285            lhs: self.lhs,
1286            rhs: self.rhs,
1287            output: self.output,
1288            rows: self.rows,
1289            columns: self.columns,
1290            contracted: self.contracted,
1291            batch_count: self.batch_count,
1292            lhs_layout: self.lhs_layout,
1293            rhs_layout: self.rhs_layout,
1294            output_layout: self.output_layout,
1295            accumulation: self.accumulation,
1296        }
1297    }
1298}
1299
1300/// Validated output-free borrowed GEMM request for full-overwrite destinations.
1301///
1302/// Mirrors [`CpuGemmRequest`] minus the writable output: the destination is
1303/// provided separately as `&mut [MaybeUninit<u8>]` to [`CpuUninitGemmProvider`]
1304/// callers. This request is only used for `beta == 0` (full-overwrite)
1305/// accumulations, where the destination is never read.
1306///
1307/// # Examples
1308///
1309/// ```
1310/// use tenferro_cpu::provider::CpuGemmUninitRequest;
1311/// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1312/// assert!(request.batch_count() >= 1);
1313/// # }
1314/// ```
1315#[derive(Debug)]
1316/// Output-free prepared GEMM request for uninitialized full-overwrite
1317/// execution (see `CpuUninitGemmProvider::gemm_into_uninit`).
1318///
1319/// # Examples
1320///
1321/// ```
1322/// use tenferro_cpu::provider::CpuGemmUninitRequest;
1323/// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1324/// #     let _ = request.rows();
1325/// #     let _ = request.columns();
1326/// #     let _ = request.contracted();
1327/// # }
1328/// ```
1329pub struct CpuGemmUninitRequest<'request, 'input> {
1330    lhs: &'request TensorRead<'input>,
1331    rhs: &'request TensorRead<'input>,
1332    rows: usize,
1333    columns: usize,
1334    contracted: usize,
1335    batch_count: usize,
1336    lhs_layout: CpuBatchedMatrixLayout,
1337    rhs_layout: CpuBatchedMatrixLayout,
1338    output_layout: CpuBatchedMatrixLayout,
1339    accumulation: DotGeneralAccumulation,
1340}
1341
1342#[cfg(any(
1343    feature = "cpu-faer",
1344    all(feature = "cpu-blas", not(feature = "provider-inject"))
1345))]
1346pub(crate) struct CpuGemmUninitRequestParts<'request, 'input> {
1347    pub(crate) lhs: &'request TensorRead<'input>,
1348    pub(crate) rhs: &'request TensorRead<'input>,
1349    pub(crate) rows: usize,
1350    pub(crate) columns: usize,
1351    pub(crate) contracted: usize,
1352    pub(crate) batch_count: usize,
1353    pub(crate) lhs_layout: CpuBatchedMatrixLayout,
1354    pub(crate) rhs_layout: CpuBatchedMatrixLayout,
1355    pub(crate) output_layout: CpuBatchedMatrixLayout,
1356    pub(crate) accumulation: DotGeneralAccumulation,
1357}
1358
1359impl<'request, 'input> CpuGemmUninitRequest<'request, 'input> {
1360    #[allow(clippy::too_many_arguments, dead_code)]
1361    pub(crate) fn new(
1362        lhs: &'request TensorRead<'input>,
1363        rhs: &'request TensorRead<'input>,
1364        rows: usize,
1365        columns: usize,
1366        contracted: usize,
1367        batch_count: usize,
1368        lhs_layout: CpuBatchedMatrixLayout,
1369        rhs_layout: CpuBatchedMatrixLayout,
1370        output_layout: CpuBatchedMatrixLayout,
1371        accumulation: DotGeneralAccumulation,
1372    ) -> Self {
1373        Self {
1374            lhs,
1375            rhs,
1376            rows,
1377            columns,
1378            contracted,
1379            batch_count,
1380            lhs_layout,
1381            rhs_layout,
1382            output_layout,
1383            accumulation,
1384        }
1385    }
1386
1387    /// Return the `lhs` of this output-free request.
1388    ///
1389    /// # Examples
1390    ///
1391    /// ```
1392    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1393    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1394    /// #     let _ = request.lhs();
1395    /// # }
1396    /// ```
1397    pub fn lhs(&self) -> &TensorRead<'input> {
1398        self.lhs
1399    }
1400
1401    /// Return the `rhs` of this output-free request.
1402    ///
1403    /// # Examples
1404    ///
1405    /// ```
1406    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1407    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1408    /// #     let _ = request.rhs();
1409    /// # }
1410    /// ```
1411    pub fn rhs(&self) -> &TensorRead<'input> {
1412        self.rhs
1413    }
1414
1415    /// Return the `rows` of this output-free request.
1416    ///
1417    /// # Examples
1418    ///
1419    /// ```
1420    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1421    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1422    /// #     let _ = request.rows();
1423    /// # }
1424    /// ```
1425    pub fn rows(&self) -> usize {
1426        self.rows
1427    }
1428
1429    /// Return the `columns` of this output-free request.
1430    ///
1431    /// # Examples
1432    ///
1433    /// ```
1434    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1435    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1436    /// #     let _ = request.columns();
1437    /// # }
1438    /// ```
1439    pub fn columns(&self) -> usize {
1440        self.columns
1441    }
1442
1443    /// Return the `contracted` of this output-free request.
1444    ///
1445    /// # Examples
1446    ///
1447    /// ```
1448    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1449    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1450    /// #     let _ = request.contracted();
1451    /// # }
1452    /// ```
1453    pub fn contracted(&self) -> usize {
1454        self.contracted
1455    }
1456
1457    /// Return the `batch_count` of this output-free request.
1458    ///
1459    /// # Examples
1460    ///
1461    /// ```
1462    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1463    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1464    /// #     let _ = request.batch_count();
1465    /// # }
1466    /// ```
1467    pub fn batch_count(&self) -> usize {
1468        self.batch_count
1469    }
1470
1471    /// Return the `lhs_layout` of this output-free request.
1472    ///
1473    /// # Examples
1474    ///
1475    /// ```
1476    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1477    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1478    /// #     let _ = request.lhs_layout();
1479    /// # }
1480    /// ```
1481    pub fn lhs_layout(&self) -> CpuBatchedMatrixLayout {
1482        self.lhs_layout
1483    }
1484
1485    /// Return the `rhs_layout` of this output-free request.
1486    ///
1487    /// # Examples
1488    ///
1489    /// ```
1490    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1491    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1492    /// #     let _ = request.rhs_layout();
1493    /// # }
1494    /// ```
1495    pub fn rhs_layout(&self) -> CpuBatchedMatrixLayout {
1496        self.rhs_layout
1497    }
1498
1499    /// Return the `output_layout` of this output-free request.
1500    ///
1501    /// # Examples
1502    ///
1503    /// ```
1504    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1505    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1506    /// #     let _ = request.output_layout();
1507    /// # }
1508    /// ```
1509    pub fn output_layout(&self) -> CpuBatchedMatrixLayout {
1510        self.output_layout
1511    }
1512
1513    /// Return the `accumulation` of this output-free request.
1514    ///
1515    /// # Examples
1516    ///
1517    /// ```
1518    /// use tenferro_cpu::provider::CpuGemmUninitRequest;
1519    /// # fn inspect(request: &CpuGemmUninitRequest<'_, '_>) {
1520    /// #     let _ = request.accumulation();
1521    /// # }
1522    /// ```
1523    pub fn accumulation(&self) -> DotGeneralAccumulation {
1524        self.accumulation
1525    }
1526
1527    #[cfg(any(
1528        feature = "cpu-faer",
1529        all(feature = "cpu-blas", not(feature = "provider-inject"))
1530    ))]
1531    pub(crate) fn into_parts(self) -> CpuGemmUninitRequestParts<'request, 'input> {
1532        CpuGemmUninitRequestParts {
1533            lhs: self.lhs,
1534            rhs: self.rhs,
1535            rows: self.rows,
1536            columns: self.columns,
1537            contracted: self.contracted,
1538            batch_count: self.batch_count,
1539            lhs_layout: self.lhs_layout,
1540            rhs_layout: self.rhs_layout,
1541            output_layout: self.output_layout,
1542            accumulation: self.accumulation,
1543        }
1544    }
1545}
1546
1547/// Validated borrowed grouped-GEMM request.
1548///
1549/// # Examples
1550///
1551/// ```
1552/// use tenferro_cpu::provider::CpuGroupedGemmRequest;
1553/// # fn inspect(request: &CpuGroupedGemmRequest<'_, '_, '_>) {
1554/// assert_eq!(request.jobs().len(), request.jobs().iter().count());
1555/// # }
1556/// ```
1557#[derive(Debug)]
1558pub struct CpuGroupedGemmRequest<'request, 'input, 'output> {
1559    lhs: &'request TensorRead<'input>,
1560    rhs: &'request TensorRead<'input>,
1561    output: &'request mut TensorWrite<'output>,
1562    jobs: &'request [GroupedGemmJob],
1563    accumulation: DotGeneralAccumulation,
1564    vendor_batch: CpuVendorBatch,
1565}
1566
1567impl<'request, 'input, 'output> CpuGroupedGemmRequest<'request, 'input, 'output> {
1568    /// This request with the engine-resolved vendor-batch control.
1569    pub(crate) fn with_vendor_batch(mut self, vendor_batch: CpuVendorBatch) -> Self {
1570        self.vendor_batch = vendor_batch;
1571        self
1572    }
1573
1574    /// Return whether the jobs may, must or must not use a vendor batch call.
1575    ///
1576    /// # Examples
1577    ///
1578    /// ```
1579    /// use tenferro_cpu::provider::{CpuGroupedGemmRequest, CpuVendorBatch};
1580    /// # fn inspect(request: &CpuGroupedGemmRequest<'_, '_, '_>) {
1581    /// let _ = request.vendor_batch() == CpuVendorBatch::Required;
1582    /// # }
1583    /// ```
1584    pub fn vendor_batch(&self) -> CpuVendorBatch {
1585        self.vendor_batch
1586    }
1587
1588    #[allow(dead_code)]
1589    pub(crate) fn new(
1590        lhs: &'request TensorRead<'input>,
1591        rhs: &'request TensorRead<'input>,
1592        output: &'request mut TensorWrite<'output>,
1593        jobs: &'request [GroupedGemmJob],
1594        accumulation: DotGeneralAccumulation,
1595    ) -> Self {
1596        Self {
1597            lhs,
1598            rhs,
1599            output,
1600            jobs,
1601            accumulation,
1602            vendor_batch: CpuVendorBatch::default(),
1603        }
1604    }
1605
1606    /// Return the borrowed left input.
1607    pub fn lhs(&self) -> &TensorRead<'input> {
1608        self.lhs
1609    }
1610
1611    /// Return the borrowed right input.
1612    pub fn rhs(&self) -> &TensorRead<'input> {
1613        self.rhs
1614    }
1615
1616    /// Reborrow the writable output.
1617    pub fn output(&mut self) -> &mut TensorWrite<'output> {
1618        self.output
1619    }
1620
1621    /// Return ordered, pairwise-disjoint validated jobs.
1622    pub fn jobs(&self) -> &[GroupedGemmJob] {
1623        self.jobs
1624    }
1625
1626    /// Return shared conjugation and alpha/beta semantics.
1627    pub fn accumulation(&self) -> DotGeneralAccumulation {
1628        self.accumulation
1629    }
1630
1631    pub(crate) fn into_parts(
1632        self,
1633    ) -> (
1634        &'request TensorRead<'input>,
1635        &'request TensorRead<'input>,
1636        &'request mut TensorWrite<'output>,
1637        &'request [GroupedGemmJob],
1638        DotGeneralAccumulation,
1639    ) {
1640        (
1641            self.lhs,
1642            self.rhs,
1643            self.output,
1644            self.jobs,
1645            self.accumulation,
1646        )
1647    }
1648}
1649
1650/// Engine-requested layout materialization.
1651///
1652/// # Examples
1653///
1654/// ```
1655/// use tenferro_cpu::provider::CpuLayoutTransformIntent;
1656/// assert_eq!(
1657///     CpuLayoutTransformIntent::CanonicalColumnMajor,
1658///     CpuLayoutTransformIntent::CanonicalColumnMajor,
1659/// );
1660/// ```
1661#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1662pub enum CpuLayoutTransformIntent {
1663    /// Materialize a compact canonical column-major tensor.
1664    CanonicalColumnMajor,
1665}
1666
1667/// Validated borrowed layout-transform request.
1668///
1669/// # Examples
1670///
1671/// ```
1672/// use tenferro_cpu::provider::{CpuLayoutTransformIntent, CpuLayoutTransformRequest};
1673/// # fn inspect(request: &CpuLayoutTransformRequest<'_, '_, '_>) {
1674/// assert_eq!(request.intent(), CpuLayoutTransformIntent::CanonicalColumnMajor);
1675/// # }
1676/// ```
1677#[derive(Debug)]
1678pub struct CpuLayoutTransformRequest<'request, 'input, 'output> {
1679    input: &'request TensorRead<'input>,
1680    output: &'request mut TensorWrite<'output>,
1681    intent: CpuLayoutTransformIntent,
1682    conjugate: bool,
1683}
1684
1685impl<'request, 'input, 'output> CpuLayoutTransformRequest<'request, 'input, 'output> {
1686    #[allow(dead_code)]
1687    pub(crate) fn new(
1688        input: &'request TensorRead<'input>,
1689        output: &'request mut TensorWrite<'output>,
1690        intent: CpuLayoutTransformIntent,
1691        conjugate: bool,
1692    ) -> Self {
1693        Self {
1694            input,
1695            output,
1696            intent,
1697            conjugate,
1698        }
1699    }
1700
1701    /// Return the borrowed input.
1702    pub fn input(&self) -> &TensorRead<'input> {
1703        self.input
1704    }
1705
1706    /// Reborrow the writable output.
1707    pub fn output(&mut self) -> &mut TensorWrite<'output> {
1708        self.output
1709    }
1710
1711    /// Return the requested materialization intent.
1712    pub fn intent(&self) -> CpuLayoutTransformIntent {
1713        self.intent
1714    }
1715
1716    /// Return whether materialization must conjugate each input element.
1717    ///
1718    /// # Examples
1719    ///
1720    /// ```
1721    /// use tenferro_cpu::provider::CpuLayoutTransformRequest;
1722    /// # fn inspect(request: &CpuLayoutTransformRequest<'_, '_, '_>) {
1723    /// let _must_conjugate = request.conjugate();
1724    /// # }
1725    /// ```
1726    pub fn conjugate(&self) -> bool {
1727        self.conjugate
1728    }
1729
1730    pub(crate) fn into_parts(
1731        self,
1732    ) -> (
1733        &'request TensorRead<'input>,
1734        &'request mut TensorWrite<'output>,
1735        CpuLayoutTransformIntent,
1736        bool,
1737    ) {
1738        (self.input, self.output, self.intent, self.conjugate)
1739    }
1740}
1741
1742/// Validated ordered contraction-role groups.
1743///
1744/// # Examples
1745///
1746/// ```
1747/// use tenferro_cpu::provider::CpuDotGeneralRequest;
1748/// # fn inspect(request: &CpuDotGeneralRequest<'_, '_, '_>) {
1749/// let _pairs = request.axes().contracting_pairs().count();
1750/// # }
1751/// ```
1752#[derive(Clone, Copy, Debug)]
1753pub struct CpuContractionAxes<'a> {
1754    lhs_rank: usize,
1755    rhs_rank: usize,
1756    lhs_contracting: &'a [usize],
1757    rhs_contracting: &'a [usize],
1758    lhs_batch: &'a [usize],
1759    rhs_batch: &'a [usize],
1760    lhs_role_mask: Option<u64>,
1761    rhs_role_mask: Option<u64>,
1762}
1763
1764impl<'a> CpuContractionAxes<'a> {
1765    #[allow(clippy::too_many_arguments, dead_code)]
1766    pub(crate) fn new(
1767        lhs_rank: usize,
1768        rhs_rank: usize,
1769        lhs_contracting: &'a [usize],
1770        rhs_contracting: &'a [usize],
1771        lhs_batch: &'a [usize],
1772        rhs_batch: &'a [usize],
1773        lhs_role_mask: Option<u64>,
1774        rhs_role_mask: Option<u64>,
1775    ) -> Self {
1776        Self {
1777            lhs_rank,
1778            rhs_rank,
1779            lhs_contracting,
1780            rhs_contracting,
1781            lhs_batch,
1782            rhs_batch,
1783            lhs_role_mask,
1784            rhs_role_mask,
1785        }
1786    }
1787
1788    /// Return ordered contracting-axis pairs.
1789    pub fn contracting_pairs(&self) -> impl ExactSizeIterator<Item = (usize, usize)> + '_ {
1790        self.lhs_contracting
1791            .iter()
1792            .copied()
1793            .zip(self.rhs_contracting.iter().copied())
1794    }
1795
1796    /// Return ordered batch-axis pairs.
1797    pub fn batch_pairs(&self) -> impl ExactSizeIterator<Item = (usize, usize)> + '_ {
1798        self.lhs_batch
1799            .iter()
1800            .copied()
1801            .zip(self.rhs_batch.iter().copied())
1802    }
1803
1804    /// Return left axes that are neither contracting nor batch axes.
1805    pub fn lhs_free_axes(&self) -> impl Iterator<Item = usize> + '_ {
1806        (0..self.lhs_rank).filter(move |&axis| !self.lhs_axis_has_role(axis))
1807    }
1808
1809    /// Return right axes that are neither contracting nor batch axes.
1810    pub fn rhs_free_axes(&self) -> impl Iterator<Item = usize> + '_ {
1811        (0..self.rhs_rank).filter(move |&axis| !self.rhs_axis_has_role(axis))
1812    }
1813
1814    fn lhs_axis_has_role(&self, axis: usize) -> bool {
1815        self.lhs_role_mask.map_or_else(
1816            || self.lhs_contracting.contains(&axis) || self.lhs_batch.contains(&axis),
1817            |mask| mask & (1_u64 << axis) != 0,
1818        )
1819    }
1820
1821    fn rhs_axis_has_role(&self, axis: usize) -> bool {
1822        self.rhs_role_mask.map_or_else(
1823            || self.rhs_contracting.contains(&axis) || self.rhs_batch.contains(&axis),
1824            |mask| mask & (1_u64 << axis) != 0,
1825        )
1826    }
1827}
1828
1829/// Validated borrowed semantic binary `dot_general` request.
1830///
1831/// # Examples
1832///
1833/// ```
1834/// use tenferro_cpu::provider::CpuDotGeneralRequest;
1835/// # fn inspect(request: &CpuDotGeneralRequest<'_, '_, '_>) {
1836/// let _ = request.accumulation();
1837/// # }
1838/// ```
1839#[derive(Debug)]
1840pub struct CpuDotGeneralRequest<'request, 'input, 'output> {
1841    lhs: &'request TensorRead<'input>,
1842    rhs: &'request TensorRead<'input>,
1843    output: &'request mut TensorWrite<'output>,
1844    axes: CpuContractionAxes<'request>,
1845    accumulation: DotGeneralAccumulation,
1846}
1847
1848impl<'request, 'input, 'output> CpuDotGeneralRequest<'request, 'input, 'output> {
1849    #[allow(dead_code)]
1850    pub(crate) fn new(
1851        lhs: &'request TensorRead<'input>,
1852        rhs: &'request TensorRead<'input>,
1853        output: &'request mut TensorWrite<'output>,
1854        axes: CpuContractionAxes<'request>,
1855        accumulation: DotGeneralAccumulation,
1856    ) -> Self {
1857        Self {
1858            lhs,
1859            rhs,
1860            output,
1861            axes,
1862            accumulation,
1863        }
1864    }
1865
1866    /// Return the borrowed left input.
1867    pub fn lhs(&self) -> &TensorRead<'input> {
1868        self.lhs
1869    }
1870
1871    /// Return the borrowed right input.
1872    pub fn rhs(&self) -> &TensorRead<'input> {
1873        self.rhs
1874    }
1875
1876    /// Reborrow the writable output.
1877    pub fn output(&mut self) -> &mut TensorWrite<'output> {
1878        self.output
1879    }
1880
1881    /// Return the validated ordered axis groups.
1882    pub fn axes(&self) -> &CpuContractionAxes<'request> {
1883        &self.axes
1884    }
1885
1886    /// Return conjugation and alpha/beta update semantics.
1887    pub fn accumulation(&self) -> DotGeneralAccumulation {
1888        self.accumulation
1889    }
1890
1891    /// Consume the request and return the validated operand borrows.
1892    ///
1893    /// External general-contraction providers use this when they need
1894    /// simultaneous immutable access to both inputs and mutable access to the
1895    /// output.
1896    ///
1897    /// # Examples
1898    ///
1899    /// ```
1900    /// use tenferro_cpu::provider::CpuDotGeneralRequest;
1901    ///
1902    /// fn consume_request(request: CpuDotGeneralRequest<'_, '_, '_>) {
1903    ///     let (_lhs, _rhs, _output, _axes, _accumulation) = request.into_parts();
1904    /// }
1905    /// ```
1906    pub fn into_parts(
1907        self,
1908    ) -> (
1909        &'request TensorRead<'input>,
1910        &'request TensorRead<'input>,
1911        &'request mut TensorWrite<'output>,
1912        CpuContractionAxes<'request>,
1913        DotGeneralAccumulation,
1914    ) {
1915        (
1916            self.lhs,
1917            self.rhs,
1918            self.output,
1919            self.axes,
1920            self.accumulation,
1921        )
1922    }
1923}
1924
1925/// Provider for validated GEMM-family requests.
1926///
1927/// # Examples
1928///
1929/// Trait objects are supported directly:
1930///
1931/// ```
1932/// use tenferro_cpu::provider::CpuGemmProvider;
1933/// # fn accepts_provider(_: &dyn CpuGemmProvider) {}
1934/// ```
1935pub trait CpuGemmProvider: fmt::Debug + Send + Sync + 'static {
1936    /// Return immutable count, placement, and fan-out capabilities.
1937    ///
1938    /// This declaration must describe controls actually applied and restored
1939    /// by the provider adapter around each call. Merely discovering a runtime
1940    /// symbol is insufficient. A provider bundle samples this method exactly
1941    /// once during construction and keeps that snapshot for its lifetime; the
1942    /// returned contract must therefore remain valid for the provider object.
1943    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
1944
1945    /// Execute one validated GEMM.
1946    ///
1947    /// # Errors
1948    ///
1949    /// Returns [`tenferro_tensor::Error::BackendSource`] or
1950    /// [`tenferro_tensor::Error::BackendFailure`] when the provider runtime
1951    /// fails. A detected inconsistency in engine-attested request metadata is
1952    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported
1953    /// capabilities use [`CpuProviderOutcome::Unsupported`] instead.
1954    fn gemm(
1955        &self,
1956        context: &CpuExecutionContext<'_>,
1957        request: CpuGemmRequest<'_, '_, '_>,
1958    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
1959
1960    /// Execute one validated strided-batched GEMM.
1961    ///
1962    /// # Errors
1963    ///
1964    /// Returns [`tenferro_tensor::Error::BackendSource`] or
1965    /// [`tenferro_tensor::Error::BackendFailure`] when the provider runtime
1966    /// fails. A detected inconsistency in engine-attested request metadata is
1967    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported
1968    /// capabilities use [`CpuProviderOutcome::Unsupported`] instead.
1969    fn strided_batched_gemm(
1970        &self,
1971        context: &CpuExecutionContext<'_>,
1972        request: CpuGemmRequest<'_, '_, '_>,
1973    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
1974
1975    /// Execute one validated grouped GEMM.
1976    ///
1977    /// # Errors
1978    ///
1979    /// Returns [`tenferro_tensor::Error::BackendSource`] or
1980    /// [`tenferro_tensor::Error::BackendFailure`] when the provider runtime
1981    /// fails. A detected inconsistency in engine-attested request metadata is
1982    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported
1983    /// capabilities use [`CpuProviderOutcome::Unsupported`] instead.
1984    fn grouped_gemm(
1985        &self,
1986        context: &CpuExecutionContext<'_>,
1987        request: CpuGroupedGemmRequest<'_, '_, '_>,
1988    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
1989
1990    /// Return the structural full-overwrite witness for this provider.
1991    ///
1992    /// Returns `Some(self)` only for types that implement
1993    /// [`CpuUninitGemmProvider`] via an `unsafe impl`, so a `Some` witness is
1994    /// structural proof that the provider asserted the full-overwrite
1995    /// destination contract. Defaults to `None` (opted out).
1996    ///
1997    /// # Examples
1998    ///
1999    /// ```
2000    /// use tenferro_cpu::provider::CpuGemmProvider;
2001    /// # fn inspect(provider: &dyn CpuGemmProvider) {
2002    /// #     let _ = provider.uninit_provider();
2003    /// # }
2004    /// ```
2005    fn uninit_provider(&self) -> Option<&dyn CpuUninitGemmProvider> {
2006        None
2007    }
2008}
2009
2010/// Provider for validated GEMM-family requests that fully overwrite their
2011/// destination (only `beta == 0` accumulations).
2012///
2013/// # Safety
2014///
2015/// Implementors assert that [`CpuUninitGemmProvider::gemm_into_uninit`] writes
2016/// every logical output element before returning
2017/// [`CpuProviderOutcome::Executed`]; partial writes make the caller's
2018/// `assume_init` undefined behavior. Implementations must never read the
2019/// destination. This trait is only invoked for `beta == 0` accumulations.
2020///
2021/// # Examples
2022///
2023/// ```
2024/// use tenferro_cpu::provider::CpuUninitGemmProvider;
2025/// # fn accepts_provider(_: &dyn CpuUninitGemmProvider) {}
2026/// ```
2027pub unsafe trait CpuUninitGemmProvider: CpuGemmProvider {
2028    /// Execute one validated full-overwrite GEMM into uninitialized bytes.
2029    ///
2030    /// # Safety
2031    ///
2032    /// Must write every element of the destination represented by
2033    /// `output_bytes` (per `request.output_layout()`) before returning
2034    /// `Executed`; partial writes are UB at the caller's `assume_init`. The
2035    /// destination is never read. Invoked only for `beta == 0` accumulations.
2036    ///
2037    /// # Errors
2038    ///
2039    /// Returns [`tenferro_tensor::Error::BackendSource`] or
2040    /// [`tenferro_tensor::Error::BackendFailure`] when the provider runtime
2041    /// fails. A detected inconsistency in engine-attested request metadata is
2042    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported
2043    /// capabilities use [`CpuProviderOutcome::Unsupported`] instead.
2044    ///
2045    /// # Examples
2046    ///
2047    /// ```
2048    /// use tenferro_cpu::provider::{CpuGemmUninitRequest, CpuProviderOutcome, CpuUninitGemmProvider};
2049    /// use tenferro_cpu::CpuExecutionContext;
2050    /// use std::mem::MaybeUninit;
2051    /// # fn inspect(
2052    /// #     provider: &dyn CpuUninitGemmProvider,
2053    /// #     context: &CpuExecutionContext<'_>,
2054    /// #     request: CpuGemmUninitRequest<'_, '_>,
2055    /// #     output_bytes: &mut [MaybeUninit<u8>],
2056    /// # ) -> Option<tenferro_tensor::Result<CpuProviderOutcome>> {
2057    /// #     let _ = (provider, context, request, output_bytes);
2058    /// #     None
2059    /// # }
2060    /// ```
2061    unsafe fn gemm_into_uninit(
2062        &self,
2063        context: &CpuExecutionContext<'_>,
2064        request: CpuGemmUninitRequest<'_, '_>,
2065        output_bytes: &mut [MaybeUninit<u8>],
2066    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
2067}
2068
2069/// Provider for engine-owned tensor materialization.
2070///
2071/// # Examples
2072///
2073/// ```
2074/// use tenferro_cpu::provider::CpuLayoutTransformProvider;
2075/// # fn accepts_provider(_: &dyn CpuLayoutTransformProvider) {}
2076/// ```
2077pub trait CpuLayoutTransformProvider: fmt::Debug + Send + Sync + 'static {
2078    /// Return immutable count, placement, and fan-out capabilities.
2079    ///
2080    /// A provider bundle samples this method exactly once during construction
2081    /// and uses the stored descriptor for all validation and dispatch.
2082    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
2083
2084    /// Materialize one validated input into a preallocated output.
2085    ///
2086    /// # Errors
2087    ///
2088    /// Returns [`tenferro_tensor::Error::BackendSource`] or
2089    /// [`tenferro_tensor::Error::BackendFailure`] when execution fails. A
2090    /// detected inconsistency in engine-attested layout or range metadata is
2091    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported layouts
2092    /// use [`CpuProviderOutcome::Unsupported`] instead.
2093    fn materialize(
2094        &self,
2095        context: &CpuExecutionContext<'_>,
2096        request: CpuLayoutTransformRequest<'_, '_, '_>,
2097    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
2098
2099    /// Return the structural full-overwrite witness for this provider.
2100    ///
2101    /// Returns `Some(self)` only for types that implement
2102    /// [`CpuUninitLayoutTransformProvider`] via an `unsafe impl`, so a `Some`
2103    /// witness is structural proof that the provider asserted the
2104    /// full-overwrite destination contract. Defaults to `None` (opted out).
2105    ///
2106    /// # Examples
2107    ///
2108    /// ```
2109    /// use tenferro_cpu::provider::CpuLayoutTransformProvider;
2110    /// # fn inspect(provider: &dyn CpuLayoutTransformProvider) {
2111    /// #     let _ = provider.uninit_provider();
2112    /// # }
2113    /// ```
2114    fn uninit_provider(&self) -> Option<&dyn CpuUninitLayoutTransformProvider> {
2115        None
2116    }
2117}
2118
2119/// Provider for engine-owned tensor materialization into uninitialized bytes.
2120///
2121/// # Safety
2122///
2123/// Implementors assert that [`CpuUninitLayoutTransformProvider::materialize_into_uninit`]
2124/// writes every element of `output_bytes` before returning
2125/// [`CpuProviderOutcome::Executed`]; partial writes make the caller's
2126/// `assume_init` undefined behavior. Implementations must never read the
2127/// destination.
2128///
2129/// # Examples
2130///
2131/// ```
2132/// use tenferro_cpu::provider::CpuUninitLayoutTransformProvider;
2133/// # fn accepts_provider(_: &dyn CpuUninitLayoutTransformProvider) {}
2134/// ```
2135pub unsafe trait CpuUninitLayoutTransformProvider: CpuLayoutTransformProvider {
2136    /// Materialize one validated input into uninitialized bytes.
2137    ///
2138    /// # Safety
2139    ///
2140    /// Must write every element of `output_bytes` (a compact column-major
2141    /// destination of the input shape for
2142    /// [`CpuLayoutTransformIntent::CanonicalColumnMajor`]) before returning
2143    /// `Executed`; partial writes are UB at the caller's `assume_init`. The
2144    /// destination is never read.
2145    ///
2146    /// # Errors
2147    ///
2148    /// Returns [`tenferro_tensor::Error::BackendSource`] or
2149    /// [`tenferro_tensor::Error::BackendFailure`] when execution fails. A
2150    /// detected inconsistency in engine-attested layout or range metadata is
2151    /// returned as [`tenferro_tensor::Error::Validation`]. Unsupported layouts
2152    /// use [`CpuProviderOutcome::Unsupported`] instead.
2153    ///
2154    /// # Examples
2155    ///
2156    /// ```
2157    /// use tenferro_cpu::provider::{CpuLayoutTransformIntent, CpuProviderOutcome, CpuUninitLayoutTransformProvider};
2158    /// use tenferro_cpu::CpuExecutionContext;
2159    /// use tenferro_tensor::TensorRead;
2160    /// use std::mem::MaybeUninit;
2161    /// # fn inspect(
2162    /// #     provider: &dyn CpuUninitLayoutTransformProvider,
2163    /// #     context: &CpuExecutionContext<'_>,
2164    /// #     input: &TensorRead<'_>,
2165    /// #     intent: CpuLayoutTransformIntent,
2166    /// #     output_bytes: &mut [MaybeUninit<u8>],
2167    /// # ) -> Option<tenferro_tensor::Result<CpuProviderOutcome>> {
2168    /// #     let _ = (provider, context, input, intent, output_bytes);
2169    /// #     None
2170    /// # }
2171    /// ```
2172    unsafe fn materialize_into_uninit(
2173        &self,
2174        context: &CpuExecutionContext<'_>,
2175        input: &TensorRead<'_>,
2176        intent: CpuLayoutTransformIntent,
2177        conjugate: bool,
2178        output_bytes: &mut [MaybeUninit<u8>],
2179    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
2180}
2181
2182/// Provider for complete validated binary `dot_general` requests.
2183///
2184/// # Examples
2185///
2186/// ```
2187/// use tenferro_cpu::provider::CpuGeneralContractionProvider;
2188/// # fn accepts_provider(_: &dyn CpuGeneralContractionProvider) {}
2189/// ```
2190pub trait CpuGeneralContractionProvider: fmt::Debug + Send + Sync + 'static {
2191    /// Return immutable count, placement, and fan-out capabilities.
2192    ///
2193    /// A provider bundle samples this method exactly once during construction
2194    /// and uses the stored descriptor for all validation and dispatch.
2195    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
2196
2197    /// Execute one complete semantic contraction.
2198    ///
2199    /// # Errors
2200    ///
2201    /// Returns [`tenferro_tensor::Error::BackendSource`] or
2202    /// [`tenferro_tensor::Error::BackendFailure`] when the provider runtime
2203    /// fails. A detected inconsistency in engine-attested axes or range
2204    /// metadata is returned as [`tenferro_tensor::Error::Validation`].
2205    /// Unsupported contractions use [`CpuProviderOutcome::Unsupported`]
2206    /// instead.
2207    fn dot_general(
2208        &self,
2209        context: &CpuExecutionContext<'_>,
2210        request: CpuDotGeneralRequest<'_, '_, '_>,
2211    ) -> tenferro_tensor::Result<CpuProviderOutcome>;
2212}
2213
2214/// Built-in faer GEMM provider.
2215///
2216/// # Examples
2217///
2218/// ```
2219/// use tenferro_cpu::provider::{CpuGemmProvider, FaerGemmProvider};
2220/// let provider: &dyn CpuGemmProvider = &FaerGemmProvider;
2221/// let _ = provider;
2222/// ```
2223#[derive(Clone, Copy, Debug, Default)]
2224pub struct FaerGemmProvider;
2225
2226impl CpuGemmProvider for FaerGemmProvider {
2227    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
2228        engine_worker_capabilities()
2229    }
2230
2231    fn gemm(
2232        &self,
2233        context: &CpuExecutionContext<'_>,
2234        request: CpuGemmRequest<'_, '_, '_>,
2235    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2236        // faer has no vendor batch routine; a forced whole-batch call is
2237        // reported instead of silently becoming a per-item loop.
2238        if request.vendor_batch() == CpuVendorBatch::Required {
2239            return Ok(CpuProviderOutcome::Unsupported(
2240                CpuProviderUnsupported::RuntimeUnavailable,
2241            ));
2242        }
2243        #[cfg(feature = "cpu-faer")]
2244        {
2245            crate::gemm::execute_faer_gemm_request(context, request)
2246        }
2247        #[cfg(not(feature = "cpu-faer"))]
2248        {
2249            let _ = (context, request);
2250            Ok(CpuProviderOutcome::Unsupported(
2251                CpuProviderUnsupported::RuntimeUnavailable,
2252            ))
2253        }
2254    }
2255
2256    fn strided_batched_gemm(
2257        &self,
2258        context: &CpuExecutionContext<'_>,
2259        request: CpuGemmRequest<'_, '_, '_>,
2260    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2261        self.gemm(context, request)
2262    }
2263
2264    fn grouped_gemm(
2265        &self,
2266        context: &CpuExecutionContext<'_>,
2267        request: CpuGroupedGemmRequest<'_, '_, '_>,
2268    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2269        if request.vendor_batch() == CpuVendorBatch::Required {
2270            return Ok(CpuProviderOutcome::Unsupported(
2271                CpuProviderUnsupported::RuntimeUnavailable,
2272            ));
2273        }
2274        #[cfg(feature = "cpu-faer")]
2275        {
2276            crate::gemm::execute_faer_grouped_request(context, request)
2277        }
2278        #[cfg(not(feature = "cpu-faer"))]
2279        {
2280            let _ = (context, request);
2281            Ok(CpuProviderOutcome::Unsupported(
2282                CpuProviderUnsupported::RuntimeUnavailable,
2283            ))
2284        }
2285    }
2286
2287    fn uninit_provider(&self) -> Option<&dyn CpuUninitGemmProvider> {
2288        Some(self)
2289    }
2290}
2291
2292// SAFETY: `gemm_into_uninit` runs the faer GEMM with `Accum::Replace` semantics
2293// (beta == 0 never reads the destination) and writes zeros for empty
2294// contractions, so every logical output element is written before `Executed`.
2295// See `crate::gemm::execute_faer_gemm_request_into_uninit`.
2296unsafe impl CpuUninitGemmProvider for FaerGemmProvider {
2297    unsafe fn gemm_into_uninit(
2298        &self,
2299        context: &CpuExecutionContext<'_>,
2300        request: CpuGemmUninitRequest<'_, '_>,
2301        output_bytes: &mut [MaybeUninit<u8>],
2302    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2303        #[cfg(feature = "cpu-faer")]
2304        {
2305            crate::gemm::execute_faer_gemm_request_into_uninit(context, request, output_bytes)
2306        }
2307        #[cfg(not(feature = "cpu-faer"))]
2308        {
2309            let _ = (context, request, output_bytes);
2310            Ok(CpuProviderOutcome::Unsupported(
2311                CpuProviderUnsupported::RuntimeUnavailable,
2312            ))
2313        }
2314    }
2315}
2316
2317/// Built-in BLAS GEMM provider.
2318///
2319/// # Examples
2320///
2321/// ```
2322/// use tenferro_cpu::provider::{BlasGemmProvider, CpuGemmProvider};
2323/// let provider: &dyn CpuGemmProvider = &BlasGemmProvider;
2324/// let _ = provider;
2325/// ```
2326#[derive(Clone, Copy, Debug, Default)]
2327pub struct BlasGemmProvider;
2328
2329impl CpuGemmProvider for BlasGemmProvider {
2330    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
2331        #[cfg(feature = "cpu-blas")]
2332        {
2333            builtin_blas_execution_capabilities()
2334        }
2335        #[cfg(not(feature = "cpu-blas"))]
2336        {
2337            // This build returns RuntimeUnavailable without invoking BLAS.
2338            serial_capabilities()
2339        }
2340    }
2341
2342    fn gemm(
2343        &self,
2344        context: &CpuExecutionContext<'_>,
2345        request: CpuGemmRequest<'_, '_, '_>,
2346    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2347        #[cfg(feature = "cpu-blas")]
2348        {
2349            crate::gemm::execute_blas_gemm_request(context, request)
2350        }
2351        #[cfg(not(feature = "cpu-blas"))]
2352        {
2353            let _ = (context, request);
2354            Ok(CpuProviderOutcome::Unsupported(
2355                CpuProviderUnsupported::RuntimeUnavailable,
2356            ))
2357        }
2358    }
2359
2360    fn strided_batched_gemm(
2361        &self,
2362        context: &CpuExecutionContext<'_>,
2363        request: CpuGemmRequest<'_, '_, '_>,
2364    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2365        self.gemm(context, request)
2366    }
2367
2368    fn grouped_gemm(
2369        &self,
2370        context: &CpuExecutionContext<'_>,
2371        request: CpuGroupedGemmRequest<'_, '_, '_>,
2372    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2373        #[cfg(feature = "cpu-blas")]
2374        {
2375            crate::gemm::execute_blas_grouped_request(context, request)
2376        }
2377        #[cfg(not(feature = "cpu-blas"))]
2378        {
2379            let _ = (context, request);
2380            Ok(CpuProviderOutcome::Unsupported(
2381                CpuProviderUnsupported::RuntimeUnavailable,
2382            ))
2383        }
2384    }
2385
2386    fn uninit_provider(&self) -> Option<&dyn CpuUninitGemmProvider> {
2387        // Injected pointers currently promise ABI compatibility, not the
2388        // stronger full-overwrite witness. Preserve their initialized path.
2389        #[cfg(feature = "provider-inject")]
2390        {
2391            None
2392        }
2393        #[cfg(not(feature = "provider-inject"))]
2394        {
2395            Some(self)
2396        }
2397    }
2398}
2399
2400// SAFETY: the implementation rejects nonzero beta, validates destination byte
2401// length/alignment, and uses BLAS's beta=0 full-overwrite contract through raw
2402// pointers. Empty contractions explicitly initialize their output to zero.
2403#[cfg(not(feature = "provider-inject"))]
2404unsafe impl CpuUninitGemmProvider for BlasGemmProvider {
2405    unsafe fn gemm_into_uninit(
2406        &self,
2407        context: &CpuExecutionContext<'_>,
2408        request: CpuGemmUninitRequest<'_, '_>,
2409        output_bytes: &mut [MaybeUninit<u8>],
2410    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2411        #[cfg(feature = "cpu-blas")]
2412        {
2413            crate::gemm::execute_blas_gemm_request_into_uninit(context, request, output_bytes)
2414        }
2415        #[cfg(not(feature = "cpu-blas"))]
2416        {
2417            let _ = (context, request, output_bytes);
2418            Ok(CpuProviderOutcome::Unsupported(
2419                CpuProviderUnsupported::RuntimeUnavailable,
2420            ))
2421        }
2422    }
2423}
2424
2425/// Built-in strided layout materialization provider.
2426///
2427/// # Examples
2428///
2429/// ```
2430/// use tenferro_cpu::provider::{CpuLayoutTransformProvider, StridedLayoutTransformProvider};
2431/// let provider: &dyn CpuLayoutTransformProvider = &StridedLayoutTransformProvider;
2432/// let _ = provider;
2433/// ```
2434#[derive(Clone, Copy, Debug, Default)]
2435pub struct StridedLayoutTransformProvider;
2436
2437impl CpuLayoutTransformProvider for StridedLayoutTransformProvider {
2438    fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
2439        engine_worker_capabilities()
2440    }
2441
2442    fn materialize(
2443        &self,
2444        context: &CpuExecutionContext<'_>,
2445        request: CpuLayoutTransformRequest<'_, '_, '_>,
2446    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2447        context.with_native_parallelism(|| materialize_strided_layout(request))
2448    }
2449
2450    fn uninit_provider(&self) -> Option<&dyn CpuUninitLayoutTransformProvider> {
2451        Some(self)
2452    }
2453}
2454
2455// SAFETY: `materialize_into_uninit` replays the layout-transform copy over the
2456// full destination (every element written via `MaybeUninit` writes from the
2457// source); zero-element destinations are trivially satisfied.
2458unsafe impl CpuUninitLayoutTransformProvider for StridedLayoutTransformProvider {
2459    unsafe fn materialize_into_uninit(
2460        &self,
2461        context: &CpuExecutionContext<'_>,
2462        input: &TensorRead<'_>,
2463        intent: CpuLayoutTransformIntent,
2464        conjugate: bool,
2465        output_bytes: &mut [MaybeUninit<u8>],
2466    ) -> tenferro_tensor::Result<CpuProviderOutcome> {
2467        context.with_native_parallelism(|| {
2468            materialize_strided_layout_into_uninit(input, intent, conjugate, output_bytes)
2469        })
2470    }
2471}
2472
2473fn materialize_strided_layout(
2474    request: CpuLayoutTransformRequest<'_, '_, '_>,
2475) -> tenferro_tensor::Result<CpuProviderOutcome> {
2476    let (input, output, _intent, conjugate) = request.into_parts();
2477    if conjugate {
2478        macro_rules! dispatch_conjugated {
2479            ($owned:ident, $view:ident) => {
2480                match (input, &mut *output) {
2481                    (
2482                        TensorRead::Tensor(input),
2483                        TensorWrite::Tensor(output),
2484                    ) if input.dtype() == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype()
2485                        && output.dtype() == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() =>
2486                    {
2487                        let input = input.as_typed::<preset_scalar!($owned)>()
2488                            .expect("the dtype guard selects this arm");
2489                        let output = output.as_typed_mut::<preset_scalar!($owned)>()
2490                            .expect("the dtype guard selects this arm");
2491                        let input = input.as_view();
2492                        let mut output = output.as_view_mut();
2493                        crate::structural::typed_conjugate_view_into(
2494                            &input,
2495                            &mut output,
2496                            "cpu layout materialization",
2497                        )?;
2498                        return Ok(CpuProviderOutcome::Executed);
2499                    }
2500                    (
2501                        TensorRead::View(TensorView::$view(input)),
2502                        TensorWrite::Tensor(output),
2503                    ) if output.dtype() == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() => {
2504                        let output = output.as_typed_mut::<preset_scalar!($owned)>()
2505                            .expect("the dtype guard selects this arm");
2506                        let mut output = output.as_view_mut();
2507                        crate::structural::typed_conjugate_view_into(
2508                            input,
2509                            &mut output,
2510                            "cpu layout materialization",
2511                        )?;
2512                        return Ok(CpuProviderOutcome::Executed);
2513                    }
2514                    (
2515                        TensorRead::Tensor(input),
2516                        TensorWrite::View(TensorViewMut::$view(output)),
2517                    ) if input.dtype() == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() => {
2518                        let input = input.as_typed::<preset_scalar!($owned)>()
2519                            .expect("the dtype guard selects this arm");
2520                        let input = input.as_view();
2521                        crate::structural::typed_conjugate_view_into(
2522                            &input,
2523                            output,
2524                            "cpu layout materialization",
2525                        )?;
2526                        return Ok(CpuProviderOutcome::Executed);
2527                    }
2528                    (
2529                        TensorRead::View(TensorView::$view(input)),
2530                        TensorWrite::View(TensorViewMut::$view(output)),
2531                    ) => {
2532                        crate::structural::typed_conjugate_view_into(
2533                            input,
2534                            output,
2535                            "cpu layout materialization",
2536                        )?;
2537                        return Ok(CpuProviderOutcome::Executed);
2538                    }
2539                    _ => {}
2540                }
2541            };
2542        }
2543        dispatch_conjugated!(F32, F32);
2544        dispatch_conjugated!(F64, F64);
2545        dispatch_conjugated!(C32, C32);
2546        dispatch_conjugated!(C64, C64);
2547        return Ok(CpuProviderOutcome::Unsupported(
2548            CpuProviderUnsupported::DType(input.dtype()),
2549        ));
2550    }
2551    macro_rules! dispatch {
2552        ($owned:ident, $view:ident) => {
2553            match (input, &mut *output) {
2554                (TensorRead::Tensor(input), TensorWrite::Tensor(output))
2555                    if input.dtype()
2556                        == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype()
2557                        && output.dtype()
2558                            == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype(
2559                            ) =>
2560                {
2561                    let input = input
2562                        .as_typed::<preset_scalar!($owned)>()
2563                        .expect("the dtype guard selects this arm");
2564                    let output = output
2565                        .as_typed_mut::<preset_scalar!($owned)>()
2566                        .expect("the dtype guard selects this arm");
2567                    let input = input.as_view();
2568                    let mut output = output.as_view_mut();
2569                    crate::structural::typed_copy_view_into(
2570                        &input,
2571                        &mut output,
2572                        "cpu layout materialization",
2573                    )?;
2574                    return Ok(CpuProviderOutcome::Executed);
2575                }
2576                (TensorRead::View(TensorView::$view(input)), TensorWrite::Tensor(output))
2577                    if output.dtype()
2578                        == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() =>
2579                {
2580                    let output = output
2581                        .as_typed_mut::<preset_scalar!($owned)>()
2582                        .expect("the dtype guard selects this arm");
2583                    let mut output = output.as_view_mut();
2584                    crate::structural::typed_copy_view_into(
2585                        input,
2586                        &mut output,
2587                        "cpu layout materialization",
2588                    )?;
2589                    return Ok(CpuProviderOutcome::Executed);
2590                }
2591                (TensorRead::Tensor(input), TensorWrite::View(TensorViewMut::$view(output)))
2592                    if input.dtype()
2593                        == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() =>
2594                {
2595                    let input = input
2596                        .as_typed::<preset_scalar!($owned)>()
2597                        .expect("the dtype guard selects this arm");
2598                    let input = input.as_view();
2599                    crate::structural::typed_copy_view_into(
2600                        &input,
2601                        output,
2602                        "cpu layout materialization",
2603                    )?;
2604                    return Ok(CpuProviderOutcome::Executed);
2605                }
2606                (
2607                    TensorRead::View(TensorView::$view(input)),
2608                    TensorWrite::View(TensorViewMut::$view(output)),
2609                ) => {
2610                    crate::structural::typed_copy_view_into(
2611                        input,
2612                        output,
2613                        "cpu layout materialization",
2614                    )?;
2615                    return Ok(CpuProviderOutcome::Executed);
2616                }
2617                _ => {}
2618            }
2619        };
2620    }
2621    dispatch!(F32, F32);
2622    dispatch!(F64, F64);
2623    dispatch!(I32, I32);
2624    dispatch!(I64, I64);
2625    dispatch!(Bool, Bool);
2626    dispatch!(C32, C32);
2627    dispatch!(C64, C64);
2628    Ok(CpuProviderOutcome::Unsupported(
2629        CpuProviderUnsupported::DType(input.dtype()),
2630    ))
2631}
2632
2633/// Full-overwrite layout transform into an uninitialized compact column-major
2634/// destination. Only [`CpuLayoutTransformIntent::CanonicalColumnMajor`] is
2635/// implemented; the destination must be exactly `element_count * size_of::<T>()`
2636/// bytes of the input dtype `T`, aligned for `T`.
2637fn materialize_strided_layout_into_uninit(
2638    input: &TensorRead<'_>,
2639    intent: CpuLayoutTransformIntent,
2640    conjugate: bool,
2641    output_bytes: &mut [MaybeUninit<u8>],
2642) -> tenferro_tensor::Result<CpuProviderOutcome> {
2643    if intent != CpuLayoutTransformIntent::CanonicalColumnMajor {
2644        return Ok(CpuProviderOutcome::Unsupported(
2645            CpuProviderUnsupported::Layout(CpuOperand::Output),
2646        ));
2647    }
2648    macro_rules! dispatch {
2649        ($owned:ident, $view:ident) => {
2650            match input {
2651                TensorRead::Tensor(input)
2652                    if input.dtype()
2653                        == <preset_scalar!($owned) as tenferro_tensor::TensorScalar>::dtype() =>
2654                {
2655                    let input = input
2656                        .as_typed::<preset_scalar!($owned)>()
2657                        .expect("the dtype guard selects this arm");
2658                    let input = input.as_view();
2659                    crate::structural::typed_copy_into_uninit(
2660                        &input,
2661                        conjugate,
2662                        output_bytes,
2663                        "cpu layout materialization",
2664                    )?;
2665                    return Ok(CpuProviderOutcome::Executed);
2666                }
2667                TensorRead::View(TensorView::$view(input)) => {
2668                    crate::structural::typed_copy_into_uninit(
2669                        input,
2670                        conjugate,
2671                        output_bytes,
2672                        "cpu layout materialization",
2673                    )?;
2674                    return Ok(CpuProviderOutcome::Executed);
2675                }
2676                _ => {}
2677            }
2678        };
2679    }
2680    dispatch!(F32, F32);
2681    dispatch!(F64, F64);
2682    dispatch!(C32, C32);
2683    dispatch!(C64, C64);
2684    Ok(CpuProviderOutcome::Unsupported(
2685        CpuProviderUnsupported::DType(input.dtype()),
2686    ))
2687}
2688
2689pub(crate) fn builtin_gemm_provider(kind: CpuBackendKind) -> Arc<dyn CpuGemmProvider> {
2690    match kind {
2691        CpuBackendKind::Faer => Arc::new(FaerGemmProvider),
2692        CpuBackendKind::Blas => Arc::new(BlasGemmProvider),
2693    }
2694}
2695
2696pub(crate) fn builtin_layout_provider() -> Arc<dyn CpuLayoutTransformProvider> {
2697    Arc::new(StridedLayoutTransformProvider)
2698}
2699
2700#[cfg(test)]
2701pub(crate) mod tests;