1use std::fmt;
2use std::mem::{size_of, size_of_val};
3use std::sync::Arc;
4
5use tenferro_runtime::program::{CoreSemanticOp, SemanticOpRef, SemanticOperationView};
6use tenferro_runtime::runtime::ImmediateEventDomainDriver;
7use tenferro_runtime::{
8 assemble_executable_engine_registration, CacheOwnerError, CoreCapabilityBundle,
9 CoreCapabilityKind, CorePrepareContext, DotGeneralPreparation, DotGeneralPrepareRequest,
10 ElementwisePrepareRequest, ElementwiseRuntime, EngineId, EngineRegistration,
11 EngineRegistrationMetadata, ExecutableEngineRegistrationConfig, ExecutionContextIdentity,
12 HardwareClassId, IndexingPrepareRequest, IndexingRuntime, InputIngressContract,
13 InputPlacementContract, InputSignature, InputSignatureContract, InputSpecializationProjection,
14 InputSpecializationRequirements, LayoutPrepareRequest, LayoutProjection, LayoutRuntime,
15 LayoutSpecialization, MemoryKind, PrepareCapability, PrepareError, PreparedOperation,
16 PreparedOperationBinding, PreparedOperationPlan, ProviderContractError, ProviderDeviceIdentity,
17 ProviderId, ReductionPrepareRequest, ReductionRuntime, ResidentOutputContract,
18 RuntimeCacheOwner, RuntimeConfigError, RuntimeInputContract, SpecializationError,
19 SpecializationProjection, SpecializationRequirements, StorageClass, TensorRead,
20 UnsupportedReason,
21};
22
23use crate::CpuBackend;
24
25const CPU_ENGINE_ID: &str = "tenferro-cpu.default.v1";
26const CPU_HARDWARE_CLASS_ID: &str = "tenferro-cpu.host.v1";
27const CPU_STORAGE_CLASS_ID: &str = "tenferro-cpu.host.v1";
28const UNKNOWN_CORE_OPERATION: &str = "unknown-core-operation";
29
30pub fn runtime_engine_id() -> Result<EngineId, RuntimeConfigError> {
37 EngineId::new(CPU_ENGINE_ID).map_err(RuntimeConfigError::from)
38}
39
40pub fn runtime_hardware_class() -> Result<HardwareClassId, RuntimeConfigError> {
47 HardwareClassId::new(CPU_HARDWARE_CLASS_ID).map_err(RuntimeConfigError::from)
48}
49
50pub fn runtime_engine_registration(
60 backend: &CpuBackend,
61) -> Result<EngineRegistration, RuntimeConfigError> {
62 runtime_engine_registration_with_id(backend, runtime_engine_id()?)
63}
64
65pub fn runtime_engine_registration_with_id(
92 backend: &CpuBackend,
93 engine_id: EngineId,
94) -> Result<EngineRegistration, RuntimeConfigError> {
95 let backend = Arc::new(backend.clone());
96 let elementwise: Arc<dyn ElementwiseRuntime> = backend.clone();
97 let reduction: Arc<dyn ReductionRuntime> = backend.clone();
98 let indexing: Arc<dyn IndexingRuntime> = backend.clone();
99 let dot_general: Arc<dyn DotGeneralPreparation> = backend.clone();
100 let layout: Arc<dyn LayoutRuntime> = backend.clone();
101 let cache_owner: Arc<dyn RuntimeCacheOwner> = backend.clone();
102 let execution_backend = backend.as_ref().clone();
103
104 let mut capabilities = CoreCapabilityBundle::builder();
105 capabilities
106 .elementwise(elementwise)
107 .reduction(reduction)
108 .indexing(indexing)
109 .dot_general(dot_general)
110 .layout(layout);
111
112 let storage = runtime_storage_class()?;
113 let default_storage = storage.clone();
114 let placement_storage = storage.clone();
115 let signature_storage = storage.clone();
116 let runtime_storage = storage.clone();
117 let resident_storage = storage.clone();
118 let allocation_domain = backend.allocation_domain();
119 let execution_info = backend.execution_info();
120 let provider_id = match execution_info.backend_kind() {
121 crate::CpuBackendKind::Faer => "tenferro.cpu.faer",
122 crate::CpuBackendKind::Blas => "tenferro.cpu.blas",
123 };
124 let provider_device_identity = ProviderDeviceIdentity::new(
125 ProviderId::new(provider_id)?,
126 format!("domain:{}", execution_info.domain_id().as_u64()),
127 )?;
128 let ingress = InputIngressContract::new(
129 InputPlacementContract::new(move |placement, candidate| {
130 candidate == &placement_storage
131 && (cpu_input_placement(placement)
132 || (placement.memory_kind == MemoryKind::Managed
133 && allocation_domain.is_some()))
134 }),
135 InputSignatureContract::new(move |placement, family, domain, candidate| {
136 candidate == &signature_storage
137 && cpu_input_signature(placement, family, domain, allocation_domain)
138 }),
139 RuntimeInputContract::new(move |input: &TensorRead<'_>, candidate| {
140 candidate == &runtime_storage && cpu_runtime_input(input, allocation_domain)
141 }),
142 ResidentOutputContract::new(move |input: &TensorRead<'_>, candidate| {
143 candidate == &resident_storage && cpu_runtime_input(input, allocation_domain)
144 }),
145 );
146 let metadata = EngineRegistrationMetadata::new(
147 engine_id,
148 provider_device_identity,
149 runtime_hardware_class()?,
150 Arc::from(vec![storage]),
151 default_storage,
152 capabilities.build(),
153 );
154 assemble_executable_engine_registration(ExecutableEngineRegistrationConfig::new(
155 metadata,
156 execution_backend,
157 Arc::new(ImmediateEventDomainDriver::new()),
158 ingress,
159 Some(cache_owner),
160 ))
161}
162
163fn cpu_input_signature(
164 placement: &tenferro_tensor::Placement,
165 backend_family: Option<&'static str>,
166 input_domain: Option<tenferro_tensor::AllocationDomainId>,
167 allocation_domain: Option<tenferro_tensor::AllocationDomainId>,
168) -> bool {
169 if placement.memory_kind == MemoryKind::Managed {
170 return backend_family.is_some()
171 && allocation_domain.is_some()
172 && input_domain == allocation_domain;
173 }
174 cpu_input_placement(placement)
175 && match backend_family {
176 None => input_domain.is_none(),
177 Some(_) => allocation_domain.is_some() && input_domain == allocation_domain,
178 }
179}
180
181fn cpu_input_placement(placement: &tenferro_tensor::Placement) -> bool {
182 matches!(
183 placement.memory_kind,
184 MemoryKind::PinnedHost | MemoryKind::UnpinnedHost
185 )
186}
187
188fn cpu_runtime_input(
189 input: &TensorRead<'_>,
190 allocation_domain: Option<tenferro_tensor::AllocationDomainId>,
191) -> bool {
192 cpu_input_signature(
193 input.placement(),
194 input.backend_family(),
195 input.allocation_domain(),
196 allocation_domain,
197 )
198}
199
200fn runtime_storage_class() -> Result<StorageClass, RuntimeConfigError> {
201 StorageClass::new(CPU_STORAGE_CLASS_ID).map_err(RuntimeConfigError::from)
202}
203
204#[derive(Clone, Copy, Debug, Eq, PartialEq)]
205enum CpuPreparedKind {
206 Elementwise,
207 Reduction,
208 Indexing,
209 DotGeneral,
210 Layout,
211}
212
213#[derive(Debug)]
214struct CpuPreparedOperation {
215 binding: PreparedOperationBinding,
216 specialization: SpecializationProjection,
217 #[allow(dead_code, reason = "bounded Debug records the selected CPU family")]
218 kind: CpuPreparedKind,
219}
220
221impl PreparedOperation for CpuPreparedOperation {
222 fn binding(&self) -> &PreparedOperationBinding {
223 &self.binding
224 }
225
226 fn specialization(&self) -> &SpecializationProjection {
227 &self.specialization
228 }
229
230 fn retained_bytes(&self) -> usize {
231 checked_specialization_heap_retained_bytes(&self.specialization).unwrap_or(usize::MAX)
232 }
233}
234
235impl ElementwiseRuntime for CpuBackend {
236 fn prepare(
237 &self,
238 request: ElementwisePrepareRequest<'_>,
239 ) -> Result<PrepareCapability, PrepareError> {
240 prepare_cpu(
241 request.operation(),
242 request.context(),
243 CpuPreparedKind::Elementwise,
244 )
245 }
246
247 fn max_fused_region_inputs(&self) -> Option<usize> {
248 Some(tenferro_cpu_fused::ERASED_FUSION_MAX_INPUTS)
249 }
250}
251
252impl ReductionRuntime for CpuBackend {
253 fn prepare(
254 &self,
255 request: ReductionPrepareRequest<'_>,
256 ) -> Result<PrepareCapability, PrepareError> {
257 prepare_cpu(
258 request.operation(),
259 request.context(),
260 CpuPreparedKind::Reduction,
261 )
262 }
263}
264
265impl IndexingRuntime for CpuBackend {
266 fn prepare(
267 &self,
268 request: IndexingPrepareRequest<'_>,
269 ) -> Result<PrepareCapability, PrepareError> {
270 prepare_cpu(
271 request.operation(),
272 request.context(),
273 CpuPreparedKind::Indexing,
274 )
275 }
276}
277
278impl DotGeneralPreparation for CpuBackend {
279 fn prepare(
280 &self,
281 request: DotGeneralPrepareRequest<'_>,
282 ) -> Result<PrepareCapability, PrepareError> {
283 prepare_cpu(
284 request.operation(),
285 request.context(),
286 CpuPreparedKind::DotGeneral,
287 )
288 }
289}
290
291impl LayoutRuntime for CpuBackend {
292 fn prepare(
293 &self,
294 request: LayoutPrepareRequest<'_>,
295 ) -> Result<PrepareCapability, PrepareError> {
296 prepare_cpu(
297 request.operation(),
298 request.context(),
299 CpuPreparedKind::Layout,
300 )
301 }
302}
303
304impl RuntimeCacheOwner for CpuBackend {
305 fn cache_stats(&self) -> Result<tenferro_runtime::runtime::CacheStats, CacheOwnerError> {
306 self.runtime_cache_stats().map_err(cache_owner_error)
307 }
308
309 fn clear_caches(&self) -> Result<(), CacheOwnerError> {
310 self.clear_runtime_caches().map_err(cache_owner_error)
311 }
312}
313
314fn prepare_cpu(
315 operation: SemanticOperationView<'_>,
316 context: &CorePrepareContext<'_>,
317 expected_kind: CpuPreparedKind,
318) -> Result<PrepareCapability, PrepareError> {
319 validate_cpu_runtime_context(context)?;
320 let SemanticOpRef::Core(op) = operation.op() else {
321 return Err(wrong_family_error(expected_kind, "extension"));
322 };
323 let Some(actual_kind) = cpu_operation_kind(op) else {
324 return Ok(PrepareCapability::Unsupported(
325 UnsupportedReason::Operation {
326 operation: UNKNOWN_CORE_OPERATION,
327 },
328 ));
329 };
330 if actual_kind != expected_kind {
331 return Err(wrong_family_error(expected_kind, core_operation_name(op)));
332 }
333
334 let minimum = minimum_specialization_requirements(actual_kind, context.inputs())?;
335 let merged =
336 merge_specialization_requirements(context.specialization().requirements(), &minimum);
337 if &merged != context.specialization().requirements() {
338 return Ok(PrepareCapability::NeedsSpecialization(merged));
339 }
340
341 Ok(PrepareCapability::Prepared(
342 PreparedOperationPlan::metadata(Arc::new(CpuPreparedOperation {
343 binding: context.binding().clone(),
344 specialization: context.specialization().clone(),
345 kind: actual_kind,
346 })),
347 ))
348}
349
350fn validate_cpu_runtime_context(context: &CorePrepareContext<'_>) -> Result<(), PrepareError> {
351 let expected_context = ExecutionContextIdentity::of::<CpuBackend>();
352 if context.binding().context_identity() != expected_context {
353 return Err(PrepareError::ProviderContract {
354 source: ProviderContractError::WrongOperationFamily {
355 expected: CoreCapabilityKind::Elementwise,
356 operation: "cpu-context-mismatch",
357 },
358 });
359 }
360 if context.binding().hardware_class().as_str() != CPU_HARDWARE_CLASS_ID {
361 return Err(PrepareError::ProviderContract {
362 source: ProviderContractError::WrongOperationFamily {
363 expected: CoreCapabilityKind::Elementwise,
364 operation: "cpu-hardware-mismatch",
365 },
366 });
367 }
368 if context.resolved_placement().storage_class().as_str() != CPU_STORAGE_CLASS_ID {
369 return Err(PrepareError::Unsupported {
370 reason: UnsupportedReason::StorageClass {
371 storage_class: context.resolved_placement().storage_class().clone(),
372 },
373 });
374 }
375 Ok(())
376}
377
378fn cpu_operation_kind(op: &CoreSemanticOp) -> Option<CpuPreparedKind> {
379 Some(match op {
380 CoreSemanticOp::Add
381 | CoreSemanticOp::Sub
382 | CoreSemanticOp::Mul
383 | CoreSemanticOp::Neg
384 | CoreSemanticOp::Conj
385 | CoreSemanticOp::Div
386 | CoreSemanticOp::Rem
387 | CoreSemanticOp::Abs
388 | CoreSemanticOp::Sign
389 | CoreSemanticOp::Maximum
390 | CoreSemanticOp::Minimum
391 | CoreSemanticOp::Compare(_)
392 | CoreSemanticOp::Select
393 | CoreSemanticOp::Clamp
394 | CoreSemanticOp::Exp
395 | CoreSemanticOp::Log
396 | CoreSemanticOp::Sin
397 | CoreSemanticOp::Cos
398 | CoreSemanticOp::Tanh
399 | CoreSemanticOp::Sqrt
400 | CoreSemanticOp::Rsqrt
401 | CoreSemanticOp::Pow
402 | CoreSemanticOp::Expm1
403 | CoreSemanticOp::Log1p
404 | CoreSemanticOp::Erf => CpuPreparedKind::Elementwise,
405 CoreSemanticOp::ReduceSum { .. }
406 | CoreSemanticOp::ReduceSumSquares { .. }
407 | CoreSemanticOp::ReduceProd { .. }
408 | CoreSemanticOp::ReduceMax { .. }
409 | CoreSemanticOp::ReduceMin { .. } => CpuPreparedKind::Reduction,
410 CoreSemanticOp::Gather(_)
411 | CoreSemanticOp::GatherDynamicSliceSizes { .. }
412 | CoreSemanticOp::Scatter(_)
413 | CoreSemanticOp::Slice(_)
414 | CoreSemanticOp::DynamicSlice { .. }
415 | CoreSemanticOp::DynamicUpdateSlice
416 | CoreSemanticOp::Pad(_)
417 | CoreSemanticOp::Concatenate { .. }
418 | CoreSemanticOp::Reverse { .. }
419 | CoreSemanticOp::ShapeOf { .. }
420 | CoreSemanticOp::DynamicTruncate { .. }
421 | CoreSemanticOp::PadToMatch { .. } => CpuPreparedKind::Indexing,
422 CoreSemanticOp::DotGeneral { .. } => CpuPreparedKind::DotGeneral,
423 CoreSemanticOp::Transpose { .. }
424 | CoreSemanticOp::Reshape { .. }
425 | CoreSemanticOp::BroadcastInDim { .. }
426 | CoreSemanticOp::Convert { .. }
427 | CoreSemanticOp::Constant { .. }
428 | CoreSemanticOp::ExtractDiag { .. }
429 | CoreSemanticOp::EmbedDiag { .. }
430 | CoreSemanticOp::Tril { .. }
431 | CoreSemanticOp::Triu { .. } => CpuPreparedKind::Layout,
432 _ => return None,
433 })
434}
435
436fn minimum_specialization_requirements(
437 kind: CpuPreparedKind,
438 inputs: &InputSignature,
439) -> Result<SpecializationRequirements, PrepareError> {
440 let mut requirements = Vec::with_capacity(inputs.entries().len());
441 for (input, entry) in inputs.entries().iter().enumerate() {
442 let mut builder = InputSpecializationRequirements::builder();
443 builder.dtype(true).rank(true);
444 match kind {
445 CpuPreparedKind::Indexing => {
446 builder.concrete_dimensions(concrete_axes_for_rank(input, entry.shape().len())?);
447 }
448 CpuPreparedKind::DotGeneral => {
449 builder
450 .concrete_dimensions(concrete_axes_for_rank(input, entry.shape().len())?)
451 .layout(LayoutSpecialization::Class);
452 }
453 CpuPreparedKind::Elementwise | CpuPreparedKind::Reduction | CpuPreparedKind::Layout => {
454 }
455 }
456 requirements.push(
457 builder
458 .build()
459 .expect("CPU minimum specialization requirements are internally valid"),
460 );
461 }
462 Ok(SpecializationRequirements::new(requirements))
463}
464
465fn concrete_axes_for_rank(input: usize, rank: usize) -> Result<Vec<u32>, PrepareError> {
466 if u32::try_from(rank).is_err() {
467 return Err(PrepareError::Specialization {
468 source: SpecializationError::ProjectionOverflow { input, rank },
469 });
470 }
471 Ok((0..rank)
472 .map(|axis| u32::try_from(axis).expect("rank precheck keeps axes encodable"))
473 .collect())
474}
475
476fn merge_specialization_requirements(
477 current: &SpecializationRequirements,
478 minimum: &SpecializationRequirements,
479) -> SpecializationRequirements {
480 debug_assert_eq!(current.inputs().len(), minimum.inputs().len());
481 let inputs = current
482 .inputs()
483 .iter()
484 .zip(minimum.inputs())
485 .map(|(current, minimum)| merge_input_requirements(current, minimum))
486 .collect::<Vec<_>>();
487 SpecializationRequirements::new(inputs)
488}
489
490fn merge_input_requirements(
491 current: &InputSpecializationRequirements,
492 minimum: &InputSpecializationRequirements,
493) -> InputSpecializationRequirements {
494 let mut axes = current.concrete_dimensions().to_vec();
495 for axis in minimum.concrete_dimensions() {
496 if !axes.contains(axis) {
497 axes.push(*axis);
498 }
499 }
500 let layout = current.layout().max(minimum.layout());
501 let rank = current.specializes_rank()
502 || minimum.specializes_rank()
503 || !axes.is_empty()
504 || layout == LayoutSpecialization::ExactStrides;
505 let alignment = match (current.alignment_log2(), minimum.alignment_log2()) {
506 (Some(left), Some(right)) => Some(left.max(right)),
507 (Some(value), None) | (None, Some(value)) => Some(value),
508 (None, None) => None,
509 };
510 let mut builder = InputSpecializationRequirements::builder();
511 builder
512 .dtype(current.specializes_dtype() || minimum.specializes_dtype())
513 .rank(rank)
514 .concrete_dimensions(axes)
515 .placement(current.placement().max(minimum.placement()))
516 .layout(layout)
517 .alignment_log2(alignment);
518 builder
519 .build()
520 .expect("merged CPU specialization requirements preserve builder invariants")
521}
522
523fn wrong_family_error(expected_kind: CpuPreparedKind, operation: &'static str) -> PrepareError {
524 PrepareError::ProviderContract {
525 source: ProviderContractError::WrongOperationFamily {
526 expected: expected_kind.core_capability(),
527 operation,
528 },
529 }
530}
531
532impl CpuPreparedKind {
533 fn core_capability(self) -> CoreCapabilityKind {
534 match self {
535 Self::Elementwise => CoreCapabilityKind::Elementwise,
536 Self::Reduction => CoreCapabilityKind::Reduction,
537 Self::Indexing => CoreCapabilityKind::Indexing,
538 Self::DotGeneral => CoreCapabilityKind::DotGeneral,
539 Self::Layout => CoreCapabilityKind::Layout,
540 }
541 }
542}
543
544fn core_operation_name(op: &CoreSemanticOp) -> &'static str {
545 match op {
546 CoreSemanticOp::Add => "add",
547 CoreSemanticOp::Sub => "sub",
548 CoreSemanticOp::Mul => "mul",
549 CoreSemanticOp::Neg => "neg",
550 CoreSemanticOp::Conj => "conj",
551 CoreSemanticOp::DotGeneral { .. } => "dot_general",
552 CoreSemanticOp::Transpose { .. } => "transpose",
553 CoreSemanticOp::Reshape { .. } => "reshape",
554 CoreSemanticOp::BroadcastInDim { .. } => "broadcast_in_dim",
555 CoreSemanticOp::Convert { .. } => "convert",
556 CoreSemanticOp::Constant { .. } => "constant",
557 CoreSemanticOp::ReduceSum { .. } => "reduce_sum",
558 CoreSemanticOp::ReduceSumSquares { .. } => "reduce_sum_squares",
559 CoreSemanticOp::Div => "div",
560 CoreSemanticOp::Rem => "rem",
561 CoreSemanticOp::Abs => "abs",
562 CoreSemanticOp::Sign => "sign",
563 CoreSemanticOp::Maximum => "maximum",
564 CoreSemanticOp::Minimum => "minimum",
565 CoreSemanticOp::Compare(_) => "compare",
566 CoreSemanticOp::Select => "select",
567 CoreSemanticOp::Clamp => "clamp",
568 CoreSemanticOp::Exp => "exp",
569 CoreSemanticOp::Log => "log",
570 CoreSemanticOp::Sin => "sin",
571 CoreSemanticOp::Cos => "cos",
572 CoreSemanticOp::Tanh => "tanh",
573 CoreSemanticOp::Sqrt => "sqrt",
574 CoreSemanticOp::Rsqrt => "rsqrt",
575 CoreSemanticOp::Pow => "pow",
576 CoreSemanticOp::Expm1 => "expm1",
577 CoreSemanticOp::Log1p => "log1p",
578 CoreSemanticOp::Erf => "erf",
579 CoreSemanticOp::ExtractDiag { .. } => "extract_diag",
580 CoreSemanticOp::EmbedDiag { .. } => "embed_diag",
581 CoreSemanticOp::Tril { .. } => "tril",
582 CoreSemanticOp::Triu { .. } => "triu",
583 CoreSemanticOp::Gather(_) => "gather",
584 CoreSemanticOp::GatherDynamicSliceSizes { .. } => "gather_dynamic_slice_sizes",
585 CoreSemanticOp::Scatter(_) => "scatter",
586 CoreSemanticOp::Slice(_) => "slice",
587 CoreSemanticOp::DynamicSlice { .. } => "dynamic_slice",
588 CoreSemanticOp::DynamicUpdateSlice => "dynamic_update_slice",
589 CoreSemanticOp::Pad(_) => "pad",
590 CoreSemanticOp::Concatenate { .. } => "concatenate",
591 CoreSemanticOp::Reverse { .. } => "reverse",
592 CoreSemanticOp::ShapeOf { .. } => "shape_of",
593 CoreSemanticOp::DynamicTruncate { .. } => "dynamic_truncate",
594 CoreSemanticOp::PadToMatch { .. } => "pad_to_match",
595 CoreSemanticOp::ReduceProd { .. } => "reduce_prod",
596 CoreSemanticOp::ReduceMax { .. } => "reduce_max",
597 CoreSemanticOp::ReduceMin { .. } => "reduce_min",
598 _ => UNKNOWN_CORE_OPERATION,
599 }
600}
601
602fn checked_specialization_heap_retained_bytes(
603 specialization: &SpecializationProjection,
604) -> Option<usize> {
605 let requirements = specialization.requirements();
606 checked_sum([
607 requirements
608 .inputs()
609 .len()
610 .checked_mul(size_of::<InputSpecializationRequirements>())?,
611 checked_sum(
612 requirements
613 .inputs()
614 .iter()
615 .map(|input| size_of_val(input.concrete_dimensions())),
616 )?,
617 specialization
618 .inputs()
619 .len()
620 .checked_mul(size_of::<InputSpecializationProjection>())?,
621 checked_sum_options(
622 specialization
623 .inputs()
624 .iter()
625 .map(input_projection_retained_bytes),
626 )?,
627 ])
628}
629
630fn input_projection_retained_bytes(projection: &InputSpecializationProjection) -> Option<usize> {
631 size_of_val(projection.concrete_dimensions()).checked_add(match projection.layout() {
632 Some(LayoutProjection::ExactStrides(strides)) if strides.spilled() => {
633 size_of_val(strides.as_slice())
634 }
635 _ => 0,
636 })
637}
638
639fn checked_sum(values: impl IntoIterator<Item = usize>) -> Option<usize> {
640 values
641 .into_iter()
642 .try_fold(0usize, |sum, value| sum.checked_add(value))
643}
644
645fn checked_sum_options(values: impl IntoIterator<Item = Option<usize>>) -> Option<usize> {
646 values
647 .into_iter()
648 .try_fold(0usize, |sum, value| sum.checked_add(value?))
649}
650
651fn cache_owner_error(source: crate::Error) -> CacheOwnerError {
652 CacheOwnerError::new(Arc::new(source))
653}
654
655impl fmt::Display for CpuPreparedKind {
656 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
657 formatter.write_str(match self {
658 Self::Elementwise => "elementwise",
659 Self::Reduction => "reduction",
660 Self::Indexing => "indexing",
661 Self::DotGeneral => "dot_general",
662 Self::Layout => "layout",
663 })
664 }
665}
666
667#[cfg(test)]
668mod tests;