Skip to main content

tenferro_cpu/
runtime_adapter.rs

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