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;