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            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;