Skip to main content

tenferro_gpu/cubecl/
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::{
7    assemble_executable_engine_registration, CacheOwnerError, CoreCapabilityBundle,
8    CoreCapabilityKind, CorePrepareContext, DotGeneralPreparation, DotGeneralPrepareRequest,
9    ElementwisePrepareRequest, ElementwiseRuntime, EngineId, EngineRegistration,
10    EngineRegistrationMetadata, ExecutableEngineRegistrationConfig, ExecutionContextIdentity,
11    HardwareClassId, IndexingPrepareRequest, IndexingRuntime, InputIngressContract,
12    InputPlacementContract, InputSignature, InputSignatureContract, InputSpecializationProjection,
13    InputSpecializationRequirements, LayoutPrepareRequest, LayoutProjection, LayoutRuntime,
14    LayoutSpecialization, PrepareCapability, PrepareError, PreparedOperation,
15    PreparedOperationBinding, PreparedOperationPlan, ProviderContractError, ProviderDeviceIdentity,
16    ProviderId, ReductionPrepareRequest, ReductionRuntime, ResidentOutputContract,
17    RuntimeCacheOwner, RuntimeConfigError, RuntimeInputContract, SpecializationError,
18    SpecializationProjection, SpecializationRequirements, StorageClass, UnsupportedReason,
19};
20use tenferro_tensor::{DeviceKind, GpuBackendKind, MemoryKind, Placement, TensorRead, TensorView};
21
22use super::event_domain::CudaEventDomainDriver;
23use super::CudaBackend;
24
25const CUDA_HARDWARE_CLASS_ID: &str = "tenferro-cuda.device.v1";
26const CUDA_STORAGE_CLASS_ID: &str = "tenferro.storage.device.v1";
27const UNKNOWN_CORE_OPERATION: &str = "unknown-core-operation";
28
29#[derive(Debug)]
30pub(crate) struct PreparedCudaRegistrationIdentity {
31    engine_id: EngineId,
32    provider_device_identity: ProviderDeviceIdentity,
33}
34
35pub(crate) fn prepare_cuda_registration_identity(
36    engine_id: EngineId,
37    device_id: super::CudaDeviceId,
38) -> Result<PreparedCudaRegistrationIdentity, RuntimeConfigError> {
39    Ok(PreparedCudaRegistrationIdentity {
40        engine_id,
41        provider_device_identity: cuda_provider_device_identity(device_id)?,
42    })
43}
44
45/// Return the canonical CUDA runtime hardware class.
46///
47/// # Errors
48///
49/// Returns [`RuntimeConfigError`] if the built-in CUDA hardware class violates
50/// runtime identifier validation.
51pub fn cuda_runtime_hardware_class() -> Result<HardwareClassId, RuntimeConfigError> {
52    HardwareClassId::new(CUDA_HARDWARE_CLASS_ID).map_err(RuntimeConfigError::from)
53}
54
55/// Build a runtime engine registration for a [`CudaBackend`] and caller-selected
56/// [`EngineId`].
57///
58/// The registration exposes CUDA direct core preparation capabilities, CUDA
59/// extension-cache ownership hooks, and the runtime-owned tensor backend
60/// execution bridge. The caller owns the engine namespace and identity; CUDA
61/// does not provide a process-global engine ID.
62///
63/// # Errors
64///
65/// Returns [`RuntimeConfigError`] if the CUDA hardware/storage identifiers or
66/// provider/device identity fail validation, or if the registration is
67/// internally invalid.
68pub fn cuda_runtime_engine_registration(
69    backend: &CudaBackend,
70    engine_id: EngineId,
71) -> Result<EngineRegistration, RuntimeConfigError> {
72    let prepared_identity = prepare_cuda_registration_identity(engine_id, backend.device_id())?;
73    let backend = Arc::new(backend.clone());
74    let elementwise: Arc<dyn ElementwiseRuntime> = backend.clone();
75    let reduction: Arc<dyn ReductionRuntime> = backend.clone();
76    let indexing: Arc<dyn IndexingRuntime> = backend.clone();
77    let dot_general: Arc<dyn DotGeneralPreparation> = backend.clone();
78    let layout: Arc<dyn LayoutRuntime> = backend.clone();
79    let cache_owner: Arc<dyn RuntimeCacheOwner> = backend.clone();
80    let execution_backend = backend.as_ref().clone();
81
82    let mut capabilities = CoreCapabilityBundle::builder();
83    capabilities
84        .elementwise(elementwise)
85        .reduction(reduction)
86        .indexing(indexing)
87        .dot_general(dot_general)
88        .layout(layout);
89
90    let storage = cuda_runtime_storage_class()?;
91    let default_storage = storage.clone();
92    let placement_storage = storage.clone();
93    let signature_storage = storage.clone();
94    let runtime_storage = storage.clone();
95    let resident_storage = storage.clone();
96    let device_ordinal = backend.device_id().ordinal() as usize;
97    let allocation_domain = backend.runtime().allocation_domain_id();
98    let ingress = InputIngressContract::new(
99        InputPlacementContract::new(move |placement, candidate| {
100            candidate == &placement_storage && cuda_input_placement(placement, device_ordinal)
101        }),
102        InputSignatureContract::new(move |placement, family, domain, candidate| {
103            candidate == &signature_storage
104                && cuda_input_signature(
105                    placement,
106                    family,
107                    domain,
108                    device_ordinal,
109                    allocation_domain,
110                )
111        }),
112        RuntimeInputContract::new(move |input: &TensorRead<'_>, candidate| {
113            candidate == &runtime_storage
114                && cuda_input_tensor(input, device_ordinal, allocation_domain)
115        }),
116        ResidentOutputContract::new(move |input: &TensorRead<'_>, candidate| {
117            candidate == &resident_storage
118                && cuda_input_tensor(input, device_ordinal, allocation_domain)
119        }),
120    );
121    let metadata = EngineRegistrationMetadata::new(
122        prepared_identity.engine_id,
123        prepared_identity.provider_device_identity,
124        cuda_runtime_hardware_class()?,
125        Arc::from(vec![storage]),
126        default_storage,
127        capabilities.build(),
128    );
129    assemble_executable_engine_registration(ExecutableEngineRegistrationConfig::new(
130        metadata,
131        execution_backend,
132        Arc::new(CudaEventDomainDriver::new(backend.runtime().clone())),
133        ingress,
134        Some(cache_owner),
135    ))
136}
137
138fn cuda_provider_device_identity(
139    device_id: super::CudaDeviceId,
140) -> Result<ProviderDeviceIdentity, RuntimeConfigError> {
141    Ok(ProviderDeviceIdentity::new(
142        ProviderId::new("tenferro.cuda")?,
143        format!("device:{}", device_id.ordinal()),
144    )?)
145}
146
147fn cuda_input_signature(
148    placement: &Placement,
149    backend_family: Option<&'static str>,
150    allocation_domain: Option<tenferro_tensor::AllocationDomainId>,
151    device_ordinal: usize,
152    expected_domain: tenferro_tensor::AllocationDomainId,
153) -> bool {
154    cuda_input_placement(placement, device_ordinal)
155        && backend_family == Some("cuda")
156        && allocation_domain == Some(expected_domain)
157}
158
159fn cuda_input_placement(placement: &Placement, device_ordinal: usize) -> bool {
160    placement.memory_kind == MemoryKind::Device
161        && matches!(
162            &placement.device,
163            Some(device)
164                if device.kind == DeviceKind::Gpu(GpuBackendKind::Cuda)
165                    && device.ordinal == device_ordinal
166        )
167}
168
169fn cuda_input_tensor(
170    input: &TensorRead<'_>,
171    device_ordinal: usize,
172    expected_domain: tenferro_tensor::AllocationDomainId,
173) -> bool {
174    // INVARIANT: TensorRead::backend_family reports the root ProviderKind
175    // (`cuda`); the CubeCL buffer's internal `cubecl` label is only a private
176    // implementation detail and is not used for runtime registration.
177    cuda_input_placement(input.placement(), device_ordinal)
178        && input.backend_family() == Some("cuda")
179        && input.allocation_domain() == Some(expected_domain)
180        && cuda_input_has_owned_buffer(input, device_ordinal, expected_domain)
181}
182
183fn cuda_input_has_owned_buffer(
184    input: &TensorRead<'_>,
185    device_ordinal: usize,
186    expected_domain: tenferro_tensor::AllocationDomainId,
187) -> bool {
188    match input.clone().tensor_view() {
189        TensorView::F32(view) => {
190            cubecl_view_has_owner::<f32>(&view, device_ordinal, expected_domain)
191        }
192        TensorView::F64(view) => {
193            cubecl_view_has_owner::<f64>(&view, device_ordinal, expected_domain)
194        }
195        TensorView::I32(view) => {
196            cubecl_view_has_owner::<i32>(&view, device_ordinal, expected_domain)
197        }
198        TensorView::I64(view) => {
199            cubecl_view_has_owner::<i64>(&view, device_ordinal, expected_domain)
200        }
201        TensorView::Bool(view) => {
202            cubecl_view_has_owner::<bool>(&view, device_ordinal, expected_domain)
203        }
204        TensorView::C32(view) => view.backend_buffer().is_some_and(|buffer| {
205            buffer
206                .as_any()
207                .downcast_ref::<crate::CubeclBuffer>()
208                .is_some_and(|buffer| {
209                    buffer.device_ordinal() == device_ordinal
210                        && buffer.allocation_domain() == expected_domain
211                })
212        }),
213        TensorView::C64(view) => view.backend_buffer().is_some_and(|buffer| {
214            buffer
215                .as_any()
216                .downcast_ref::<crate::CubeclBuffer>()
217                .is_some_and(|buffer| {
218                    buffer.device_ordinal() == device_ordinal
219                        && buffer.allocation_domain() == expected_domain
220                })
221        }),
222    }
223}
224
225fn cubecl_view_has_owner<T: 'static>(
226    view: &tenferro_tensor::TypedTensorView<'_, T>,
227    device_ordinal: usize,
228    expected_domain: tenferro_tensor::AllocationDomainId,
229) -> bool {
230    view.backend_buffer().is_some_and(|buffer| {
231        buffer
232            .as_any()
233            .downcast_ref::<crate::CubeclBuffer>()
234            .is_some_and(|buffer| {
235                buffer.device_ordinal() == device_ordinal
236                    && buffer.allocation_domain() == expected_domain
237            })
238    })
239}
240
241#[cfg(test)]
242#[path = "tests/runtime_adapter.rs"]
243mod tests;
244
245fn cuda_runtime_storage_class() -> Result<StorageClass, RuntimeConfigError> {
246    StorageClass::new(CUDA_STORAGE_CLASS_ID).map_err(RuntimeConfigError::from)
247}
248
249#[derive(Clone, Copy, Debug, Eq, PartialEq)]
250enum CudaPreparedKind {
251    Elementwise,
252    Reduction,
253    Indexing,
254    DotGeneral,
255    Layout,
256}
257
258#[derive(Debug)]
259struct CudaPreparedOperation {
260    binding: PreparedOperationBinding,
261    specialization: SpecializationProjection,
262    #[allow(dead_code, reason = "bounded Debug records the selected CUDA family")]
263    kind: CudaPreparedKind,
264}
265
266impl PreparedOperation for CudaPreparedOperation {
267    fn binding(&self) -> &PreparedOperationBinding {
268        &self.binding
269    }
270
271    fn specialization(&self) -> &SpecializationProjection {
272        &self.specialization
273    }
274
275    fn retained_bytes(&self) -> usize {
276        checked_specialization_heap_retained_bytes(&self.specialization).unwrap_or(usize::MAX)
277    }
278}
279
280impl ElementwiseRuntime for CudaBackend {
281    fn prepare(
282        &self,
283        request: ElementwisePrepareRequest<'_>,
284    ) -> Result<PrepareCapability, PrepareError> {
285        prepare_cuda(
286            request.operation(),
287            request.context(),
288            CudaPreparedKind::Elementwise,
289        )
290    }
291}
292
293impl ReductionRuntime for CudaBackend {
294    fn prepare(
295        &self,
296        request: ReductionPrepareRequest<'_>,
297    ) -> Result<PrepareCapability, PrepareError> {
298        prepare_cuda(
299            request.operation(),
300            request.context(),
301            CudaPreparedKind::Reduction,
302        )
303    }
304}
305
306impl IndexingRuntime for CudaBackend {
307    fn prepare(
308        &self,
309        request: IndexingPrepareRequest<'_>,
310    ) -> Result<PrepareCapability, PrepareError> {
311        prepare_cuda(
312            request.operation(),
313            request.context(),
314            CudaPreparedKind::Indexing,
315        )
316    }
317}
318
319impl DotGeneralPreparation for CudaBackend {
320    fn prepare(
321        &self,
322        request: DotGeneralPrepareRequest<'_>,
323    ) -> Result<PrepareCapability, PrepareError> {
324        prepare_cuda(
325            request.operation(),
326            request.context(),
327            CudaPreparedKind::DotGeneral,
328        )
329    }
330}
331
332impl LayoutRuntime for CudaBackend {
333    fn prepare(
334        &self,
335        request: LayoutPrepareRequest<'_>,
336    ) -> Result<PrepareCapability, PrepareError> {
337        prepare_cuda(
338            request.operation(),
339            request.context(),
340            CudaPreparedKind::Layout,
341        )
342    }
343}
344
345impl RuntimeCacheOwner for CudaBackend {
346    fn cache_stats(&self) -> Result<tenferro_runtime::runtime::CacheStats, CacheOwnerError> {
347        let stats = self
348            .cuda_extension_cache_stats()
349            .map_err(cache_owner_error)?;
350        Ok(tenferro_runtime::runtime::CacheStats {
351            entries: stats.entries,
352            retained_bytes: stats.retained_bytes,
353            hits: stats.hits,
354            misses: stats.misses,
355            evictions: stats.evictions,
356            clears: stats.clears,
357        })
358    }
359
360    fn clear_caches(&self) -> Result<(), CacheOwnerError> {
361        self.clear_cuda_extension_cache().map_err(cache_owner_error)
362    }
363}
364
365fn prepare_cuda(
366    operation: SemanticOperationView<'_>,
367    context: &CorePrepareContext<'_>,
368    expected_kind: CudaPreparedKind,
369) -> Result<PrepareCapability, PrepareError> {
370    validate_cuda_runtime_context(context)?;
371    let SemanticOpRef::Core(op) = operation.op() else {
372        return Err(wrong_family_error(expected_kind, "extension"));
373    };
374    let Some(actual_kind) = cuda_operation_kind(op) else {
375        return Ok(PrepareCapability::Unsupported(
376            UnsupportedReason::Operation {
377                operation: UNKNOWN_CORE_OPERATION,
378            },
379        ));
380    };
381    if actual_kind != expected_kind {
382        return Err(wrong_family_error(expected_kind, core_operation_name(op)));
383    }
384
385    let minimum = minimum_specialization_requirements(actual_kind, context.inputs())?;
386    let merged =
387        merge_specialization_requirements(context.specialization().requirements(), &minimum);
388    if &merged != context.specialization().requirements() {
389        return Ok(PrepareCapability::NeedsSpecialization(merged));
390    }
391
392    Ok(PrepareCapability::Prepared(
393        PreparedOperationPlan::metadata(Arc::new(CudaPreparedOperation {
394            binding: context.binding().clone(),
395            specialization: context.specialization().clone(),
396            kind: actual_kind,
397        })),
398    ))
399}
400
401fn validate_cuda_runtime_context(context: &CorePrepareContext<'_>) -> Result<(), PrepareError> {
402    let expected_context = ExecutionContextIdentity::of::<CudaBackend>();
403    if context.binding().context_identity() != expected_context {
404        return Err(PrepareError::ProviderContract {
405            source: ProviderContractError::WrongOperationFamily {
406                expected: CoreCapabilityKind::Elementwise,
407                operation: "cuda-context-mismatch",
408            },
409        });
410    }
411    if context.binding().hardware_class().as_str() != CUDA_HARDWARE_CLASS_ID {
412        return Err(PrepareError::ProviderContract {
413            source: ProviderContractError::WrongOperationFamily {
414                expected: CoreCapabilityKind::Elementwise,
415                operation: "cuda-hardware-mismatch",
416            },
417        });
418    }
419    if context.resolved_placement().storage_class().as_str() != CUDA_STORAGE_CLASS_ID {
420        return Err(PrepareError::Unsupported {
421            reason: UnsupportedReason::StorageClass {
422                storage_class: context.resolved_placement().storage_class().clone(),
423            },
424        });
425    }
426    Ok(())
427}
428
429fn cuda_operation_kind(op: &CoreSemanticOp) -> Option<CudaPreparedKind> {
430    Some(match op {
431        CoreSemanticOp::Add
432        | CoreSemanticOp::Sub
433        | CoreSemanticOp::Mul
434        | CoreSemanticOp::Neg
435        | CoreSemanticOp::Conj
436        | CoreSemanticOp::Div
437        | CoreSemanticOp::Rem
438        | CoreSemanticOp::Abs
439        | CoreSemanticOp::Sign
440        | CoreSemanticOp::Maximum
441        | CoreSemanticOp::Minimum
442        | CoreSemanticOp::Compare(_)
443        | CoreSemanticOp::Select
444        | CoreSemanticOp::Clamp
445        | CoreSemanticOp::Exp
446        | CoreSemanticOp::Log
447        | CoreSemanticOp::Sin
448        | CoreSemanticOp::Cos
449        | CoreSemanticOp::Tanh
450        | CoreSemanticOp::Sqrt
451        | CoreSemanticOp::Rsqrt
452        | CoreSemanticOp::Pow
453        | CoreSemanticOp::Expm1
454        | CoreSemanticOp::Log1p
455        | CoreSemanticOp::Erf => CudaPreparedKind::Elementwise,
456        CoreSemanticOp::ReduceSum { .. }
457        | CoreSemanticOp::ReduceSumSquares { .. }
458        | CoreSemanticOp::ReduceProd { .. }
459        | CoreSemanticOp::ReduceMax { .. }
460        | CoreSemanticOp::ReduceMin { .. } => CudaPreparedKind::Reduction,
461        CoreSemanticOp::Gather(_)
462        | CoreSemanticOp::GatherDynamicSliceSizes { .. }
463        | CoreSemanticOp::Scatter(_)
464        | CoreSemanticOp::Slice(_)
465        | CoreSemanticOp::DynamicSlice { .. }
466        | CoreSemanticOp::DynamicUpdateSlice
467        | CoreSemanticOp::Pad(_)
468        | CoreSemanticOp::Concatenate { .. }
469        | CoreSemanticOp::Reverse { .. }
470        | CoreSemanticOp::ShapeOf { .. }
471        | CoreSemanticOp::DynamicTruncate { .. }
472        | CoreSemanticOp::PadToMatch { .. } => CudaPreparedKind::Indexing,
473        CoreSemanticOp::DotGeneral { .. } => CudaPreparedKind::DotGeneral,
474        CoreSemanticOp::Transpose { .. }
475        | CoreSemanticOp::Reshape { .. }
476        | CoreSemanticOp::BroadcastInDim { .. }
477        | CoreSemanticOp::Convert { .. }
478        | CoreSemanticOp::Constant { .. }
479        | CoreSemanticOp::ExtractDiag { .. }
480        | CoreSemanticOp::EmbedDiag { .. }
481        | CoreSemanticOp::Tril { .. }
482        | CoreSemanticOp::Triu { .. } => CudaPreparedKind::Layout,
483        _ => return None,
484    })
485}
486
487fn minimum_specialization_requirements(
488    kind: CudaPreparedKind,
489    inputs: &InputSignature,
490) -> Result<SpecializationRequirements, PrepareError> {
491    let mut requirements = Vec::with_capacity(inputs.entries().len());
492    for (input, entry) in inputs.entries().iter().enumerate() {
493        let mut builder = InputSpecializationRequirements::builder();
494        builder.dtype(true).rank(true);
495        match kind {
496            CudaPreparedKind::Indexing => {
497                builder.concrete_dimensions(concrete_axes_for_rank(input, entry.shape().len())?);
498            }
499            CudaPreparedKind::DotGeneral => {
500                builder
501                    .concrete_dimensions(concrete_axes_for_rank(input, entry.shape().len())?)
502                    .layout(LayoutSpecialization::Class);
503            }
504            CudaPreparedKind::Elementwise
505            | CudaPreparedKind::Reduction
506            | CudaPreparedKind::Layout => {}
507        }
508        requirements.push(
509            builder
510                .build()
511                .expect("CUDA minimum specialization requirements are internally valid"),
512        );
513    }
514    Ok(SpecializationRequirements::new(requirements))
515}
516
517fn concrete_axes_for_rank(input: usize, rank: usize) -> Result<Vec<u32>, PrepareError> {
518    if u32::try_from(rank).is_err() {
519        return Err(PrepareError::Specialization {
520            source: SpecializationError::ProjectionOverflow { input, rank },
521        });
522    }
523    Ok((0..rank)
524        .map(|axis| u32::try_from(axis).expect("rank precheck keeps axes encodable"))
525        .collect())
526}
527
528fn merge_specialization_requirements(
529    current: &SpecializationRequirements,
530    minimum: &SpecializationRequirements,
531) -> SpecializationRequirements {
532    debug_assert_eq!(current.inputs().len(), minimum.inputs().len());
533    let inputs = current
534        .inputs()
535        .iter()
536        .zip(minimum.inputs())
537        .map(|(current, minimum)| merge_input_requirements(current, minimum))
538        .collect::<Vec<_>>();
539    SpecializationRequirements::new(inputs)
540}
541
542fn merge_input_requirements(
543    current: &InputSpecializationRequirements,
544    minimum: &InputSpecializationRequirements,
545) -> InputSpecializationRequirements {
546    let mut axes = current.concrete_dimensions().to_vec();
547    for axis in minimum.concrete_dimensions() {
548        if !axes.contains(axis) {
549            axes.push(*axis);
550        }
551    }
552    let layout = current.layout().max(minimum.layout());
553    let rank = current.specializes_rank()
554        || minimum.specializes_rank()
555        || !axes.is_empty()
556        || layout == LayoutSpecialization::ExactStrides;
557    let alignment = match (current.alignment_log2(), minimum.alignment_log2()) {
558        (Some(left), Some(right)) => Some(left.max(right)),
559        (Some(value), None) | (None, Some(value)) => Some(value),
560        (None, None) => None,
561    };
562    let mut builder = InputSpecializationRequirements::builder();
563    builder
564        .dtype(current.specializes_dtype() || minimum.specializes_dtype())
565        .rank(rank)
566        .concrete_dimensions(axes)
567        .placement(current.placement().max(minimum.placement()))
568        .layout(layout)
569        .alignment_log2(alignment);
570    builder
571        .build()
572        .expect("merged CUDA specialization requirements preserve builder invariants")
573}
574
575fn wrong_family_error(expected_kind: CudaPreparedKind, operation: &'static str) -> PrepareError {
576    PrepareError::ProviderContract {
577        source: ProviderContractError::WrongOperationFamily {
578            expected: expected_kind.core_capability(),
579            operation,
580        },
581    }
582}
583
584impl CudaPreparedKind {
585    fn core_capability(self) -> CoreCapabilityKind {
586        match self {
587            Self::Elementwise => CoreCapabilityKind::Elementwise,
588            Self::Reduction => CoreCapabilityKind::Reduction,
589            Self::Indexing => CoreCapabilityKind::Indexing,
590            Self::DotGeneral => CoreCapabilityKind::DotGeneral,
591            Self::Layout => CoreCapabilityKind::Layout,
592        }
593    }
594}
595
596fn core_operation_name(op: &CoreSemanticOp) -> &'static str {
597    match op {
598        CoreSemanticOp::Add => "add",
599        CoreSemanticOp::Sub => "sub",
600        CoreSemanticOp::Mul => "mul",
601        CoreSemanticOp::Neg => "neg",
602        CoreSemanticOp::Conj => "conj",
603        CoreSemanticOp::DotGeneral { .. } => "dot_general",
604        CoreSemanticOp::Transpose { .. } => "transpose",
605        CoreSemanticOp::Reshape { .. } => "reshape",
606        CoreSemanticOp::BroadcastInDim { .. } => "broadcast_in_dim",
607        CoreSemanticOp::Convert { .. } => "convert",
608        CoreSemanticOp::Constant { .. } => "constant",
609        CoreSemanticOp::ReduceSum { .. } => "reduce_sum",
610        CoreSemanticOp::ReduceSumSquares { .. } => "reduce_sum_squares",
611        CoreSemanticOp::Div => "div",
612        CoreSemanticOp::Rem => "rem",
613        CoreSemanticOp::Abs => "abs",
614        CoreSemanticOp::Sign => "sign",
615        CoreSemanticOp::Maximum => "maximum",
616        CoreSemanticOp::Minimum => "minimum",
617        CoreSemanticOp::Compare(_) => "compare",
618        CoreSemanticOp::Select => "select",
619        CoreSemanticOp::Clamp => "clamp",
620        CoreSemanticOp::Exp => "exp",
621        CoreSemanticOp::Log => "log",
622        CoreSemanticOp::Sin => "sin",
623        CoreSemanticOp::Cos => "cos",
624        CoreSemanticOp::Tanh => "tanh",
625        CoreSemanticOp::Sqrt => "sqrt",
626        CoreSemanticOp::Rsqrt => "rsqrt",
627        CoreSemanticOp::Pow => "pow",
628        CoreSemanticOp::Expm1 => "expm1",
629        CoreSemanticOp::Log1p => "log1p",
630        CoreSemanticOp::Erf => "erf",
631        CoreSemanticOp::ExtractDiag { .. } => "extract_diag",
632        CoreSemanticOp::EmbedDiag { .. } => "embed_diag",
633        CoreSemanticOp::Tril { .. } => "tril",
634        CoreSemanticOp::Triu { .. } => "triu",
635        CoreSemanticOp::Gather(_) => "gather",
636        CoreSemanticOp::GatherDynamicSliceSizes { .. } => "gather_dynamic_slice_sizes",
637        CoreSemanticOp::Scatter(_) => "scatter",
638        CoreSemanticOp::Slice(_) => "slice",
639        CoreSemanticOp::DynamicSlice { .. } => "dynamic_slice",
640        CoreSemanticOp::DynamicUpdateSlice => "dynamic_update_slice",
641        CoreSemanticOp::Pad(_) => "pad",
642        CoreSemanticOp::Concatenate { .. } => "concatenate",
643        CoreSemanticOp::Reverse { .. } => "reverse",
644        CoreSemanticOp::ShapeOf { .. } => "shape_of",
645        CoreSemanticOp::DynamicTruncate { .. } => "dynamic_truncate",
646        CoreSemanticOp::PadToMatch { .. } => "pad_to_match",
647        CoreSemanticOp::ReduceProd { .. } => "reduce_prod",
648        CoreSemanticOp::ReduceMax { .. } => "reduce_max",
649        CoreSemanticOp::ReduceMin { .. } => "reduce_min",
650        _ => UNKNOWN_CORE_OPERATION,
651    }
652}
653
654fn checked_specialization_heap_retained_bytes(
655    specialization: &SpecializationProjection,
656) -> Option<usize> {
657    let requirements = specialization.requirements();
658    checked_sum([
659        requirements
660            .inputs()
661            .len()
662            .checked_mul(size_of::<InputSpecializationRequirements>())?,
663        checked_sum(
664            requirements
665                .inputs()
666                .iter()
667                .map(|input| size_of_val(input.concrete_dimensions())),
668        )?,
669        specialization
670            .inputs()
671            .len()
672            .checked_mul(size_of::<InputSpecializationProjection>())?,
673        checked_sum_options(
674            specialization
675                .inputs()
676                .iter()
677                .map(input_projection_retained_bytes),
678        )?,
679    ])
680}
681
682fn input_projection_retained_bytes(projection: &InputSpecializationProjection) -> Option<usize> {
683    size_of_val(projection.concrete_dimensions()).checked_add(match projection.layout() {
684        Some(LayoutProjection::ExactStrides(strides)) if strides.spilled() => {
685            size_of_val(strides.as_slice())
686        }
687        _ => 0,
688    })
689}
690
691fn checked_sum(values: impl IntoIterator<Item = usize>) -> Option<usize> {
692    values
693        .into_iter()
694        .try_fold(0usize, |sum, value| sum.checked_add(value))
695}
696
697fn checked_sum_options(values: impl IntoIterator<Item = Option<usize>>) -> Option<usize> {
698    values
699        .into_iter()
700        .try_fold(0usize, |sum, value| sum.checked_add(value?))
701}
702
703fn cache_owner_error(source: crate::Error) -> CacheOwnerError {
704    CacheOwnerError::new(Arc::new(source))
705}
706
707impl fmt::Display for CudaPreparedKind {
708    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
709        formatter.write_str(match self {
710            Self::Elementwise => "elementwise",
711            Self::Reduction => "reduction",
712            Self::Indexing => "indexing",
713            Self::DotGeneral => "dot_general",
714            Self::Layout => "layout",
715        })
716    }
717}