Skip to main content

laddu_runtime/cpu/
cache.rs

1#[cfg(feature = "jit")]
2use std::marker::PhantomData;
3use std::{mem::size_of, sync::Arc};
4
5use laddu_compile::CachePlan;
6use laddu_data::{
7    data::{CacheStorage, Dataset, EventBatch},
8    io::ReadPlan,
9};
10use laddu_expr::{ExprId, ExprNode, P4Component, ValueKind};
11use nalgebra::{DMatrix, DVector};
12use num::complex::Complex64;
13
14use super::evaluation::EventColumn;
15use super::layout::{FlatRows, eval_binary, eval_unary, matrix_at_optional};
16use super::{CpuPlan, DynamicLu, PreparedDatasetStats, RuntimeError, RuntimeResult, Value};
17use crate::MemoryLease;
18
19/// Raw cache payload metadata consumed by compiled JIT kernels.
20#[cfg(feature = "jit")]
21#[repr(C)]
22#[derive(Copy, Clone)]
23pub(crate) struct CacheDescriptor {
24    pub(crate) values: *const u8,
25    pub(crate) width: usize,
26}
27
28#[cfg(feature = "jit")]
29pub(crate) struct JitDescriptorSet<'a> {
30    pub(crate) values: Vec<CacheDescriptor>,
31    pub(crate) solve_rows: Vec<CacheDescriptor>,
32    pub(crate) _cache: PhantomData<&'a CpuBatchCache>,
33}
34
35/// Materialized event-dependent values for one batch.
36#[derive(Clone, Debug)]
37pub struct CpuBatchCache {
38    pub(super) len: usize,
39    pub(super) weights: Vec<f64>,
40    pub(super) sum_weights: f64,
41    pub(super) nodes: Vec<ExprId>,
42    pub(crate) slots: Vec<CachedSlot>,
43    pub(super) factor_nodes: Vec<ExprId>,
44    pub(super) factor_slots: Vec<CachedFactorSlot>,
45    pub(super) solve_row_keys: Vec<(ExprId, usize, usize)>,
46    pub(crate) solve_row_slots: Vec<CachedSolveRowSlot>,
47}
48
49impl CpuBatchCache {
50    pub(super) fn new(
51        cache_plan: &CachePlan,
52        factor_matrices: &[(ExprId, usize)],
53        solve_row_keys: &[(ExprId, usize, usize)],
54        len: usize,
55    ) -> RuntimeResult<Self> {
56        let slots = cache_plan
57            .entries()
58            .iter()
59            .map(|entry| CachedSlot::new(entry.value_kind(), len))
60            .collect::<RuntimeResult<Vec<_>>>()?;
61        let solve_row_slots = solve_row_keys
62            .iter()
63            .map(|(_, _, dimension)| CachedSolveRowSlot::new(*dimension, len))
64            .collect::<RuntimeResult<Vec<_>>>()?;
65        Ok(Self {
66            len,
67            weights: vec![1.0; len],
68            sum_weights: len as f64,
69            nodes: cache_plan
70                .entries()
71                .iter()
72                .map(|entry| entry.node())
73                .collect(),
74            slots,
75            factor_nodes: factor_matrices.iter().map(|(node, _)| *node).collect(),
76            factor_slots: factor_matrices
77                .iter()
78                .map(|(_, dimension)| CachedFactorSlot::new(*dimension))
79                .collect(),
80            solve_row_keys: solve_row_keys.to_vec(),
81            solve_row_slots,
82        })
83    }
84
85    /// Returns the number of cached events.
86    pub fn len(&self) -> usize {
87        self.len
88    }
89
90    /// Returns whether the cache contains no events.
91    pub fn is_empty(&self) -> bool {
92        self.len == 0
93    }
94
95    /// Returns per-event weights.
96    pub fn weights(&self) -> &[f64] {
97        &self.weights
98    }
99
100    /// Returns the sum of event weights.
101    pub fn sum_weights(&self) -> f64 {
102        self.sum_weights
103    }
104
105    /// Returns the raw cache payload metadata required by the JIT ABI.
106    ///
107    /// The returned pointers borrow this cache and are valid only while it is
108    /// immutably borrowed.  Keeping this projection here prevents the JIT
109    /// backend from depending on cache-slot representation details.
110    #[cfg(feature = "jit")]
111    #[cfg(feature = "jit")]
112    pub(crate) fn jit_descriptors(&self) -> JitDescriptorSet<'_> {
113        JitDescriptorSet {
114            values: self
115                .slots
116                .iter()
117                .map(|slot| CacheDescriptor {
118                    values: slot.values_ptr(),
119                    width: slot.width(),
120                })
121                .collect(),
122            solve_rows: self
123                .solve_row_slots
124                .iter()
125                .map(|slot| CacheDescriptor {
126                    values: slot.values.as_ptr().cast(),
127                    width: slot.dimension,
128                })
129                .collect(),
130            _cache: PhantomData,
131        }
132    }
133
134    /// Estimates heap memory retained by this cache, in bytes.
135    pub fn resident_bytes(&self) -> usize {
136        self.weights.capacity() * size_of::<f64>()
137            + self.nodes.capacity() * size_of::<ExprId>()
138            + self
139                .slots
140                .iter()
141                .map(CachedSlot::resident_bytes)
142                .sum::<usize>()
143            + self.factor_nodes.capacity() * size_of::<ExprId>()
144            + self
145                .factor_slots
146                .iter()
147                .map(CachedFactorSlot::resident_bytes)
148                .sum::<usize>()
149            + self.solve_row_keys.capacity() * size_of::<(ExprId, usize, usize)>()
150            + self
151                .solve_row_slots
152                .iter()
153                .map(CachedSolveRowSlot::resident_bytes)
154                .sum::<usize>()
155    }
156
157    pub(super) fn set_weights(&mut self, weights: Vec<f64>) {
158        self.sum_weights = weights.iter().sum();
159        self.weights = weights;
160    }
161
162    pub(super) fn push(&mut self, slot: usize, value: Value) -> RuntimeResult<()> {
163        let len = self.slots.len();
164        self.slots
165            .get_mut(slot)
166            .ok_or(RuntimeError::InvalidCache {
167                expected: len,
168                actual: slot + 1,
169            })?
170            .push(value)
171    }
172
173    pub(super) fn value(&self, slot: usize, row: usize) -> RuntimeResult<Value> {
174        if row >= self.len {
175            return Err(RuntimeError::InvalidShape {
176                index: row,
177                message: format!("cache row {row} out of bounds for len {}", self.len),
178            });
179        }
180        self.slots
181            .get(slot)
182            .ok_or(RuntimeError::InvalidCache {
183                expected: self.slots.len(),
184                actual: slot + 1,
185            })?
186            .value(row)
187    }
188
189    pub(super) fn scalar(&self, slot: usize, row: usize) -> RuntimeResult<Complex64> {
190        if row >= self.len {
191            return Err(RuntimeError::InvalidShape {
192                index: row,
193                message: format!("cache row {row} out of bounds for len {}", self.len),
194            });
195        }
196        self.slots
197            .get(slot)
198            .ok_or(RuntimeError::InvalidCache {
199                expected: self.slots.len(),
200                actual: slot + 1,
201            })?
202            .scalar(row)
203    }
204
205    pub(super) fn real_range(
206        &self,
207        slot: usize,
208        start: usize,
209        end: usize,
210    ) -> RuntimeResult<&[f64]> {
211        if start > end || end > self.len {
212            return Err(RuntimeError::InvalidShape {
213                index: start,
214                message: format!(
215                    "cache range {start}..{end} out of bounds for len {}",
216                    self.len
217                ),
218            });
219        }
220        self.slots
221            .get(slot)
222            .ok_or(RuntimeError::InvalidCache {
223                expected: self.slots.len(),
224                actual: slot + 1,
225            })?
226            .real_range(start, end)
227    }
228
229    pub(super) fn complex_range(
230        &self,
231        slot: usize,
232        start: usize,
233        end: usize,
234    ) -> RuntimeResult<&[Complex64]> {
235        if start > end || end > self.len {
236            return Err(RuntimeError::InvalidShape {
237                index: start,
238                message: format!(
239                    "cache range {start}..{end} out of bounds for len {}",
240                    self.len
241                ),
242            });
243        }
244        self.slots
245            .get(slot)
246            .ok_or(RuntimeError::InvalidCache {
247                expected: self.slots.len(),
248                actual: slot + 1,
249            })?
250            .complex_range(start, end)
251    }
252
253    pub(super) fn push_factor(&mut self, slot: usize, factor: DynamicLu) -> RuntimeResult<()> {
254        let len = self.factor_slots.len();
255        self.factor_slots
256            .get_mut(slot)
257            .ok_or(RuntimeError::InvalidCache {
258                expected: len,
259                actual: slot + 1,
260            })?
261            .push(factor)
262    }
263
264    pub(super) fn factor(&self, slot: usize, row: usize) -> RuntimeResult<&DynamicLu> {
265        self.factor_slots
266            .get(slot)
267            .ok_or(RuntimeError::InvalidCache {
268                expected: self.factor_slots.len(),
269                actual: slot + 1,
270            })?
271            .factor(row)
272    }
273
274    pub(super) fn push_solve_row(
275        &mut self,
276        slot: usize,
277        values: impl IntoIterator<Item = Complex64>,
278    ) -> RuntimeResult<()> {
279        let len = self.solve_row_slots.len();
280        self.solve_row_slots
281            .get_mut(slot)
282            .ok_or(RuntimeError::InvalidCache {
283                expected: len,
284                actual: slot + 1,
285            })?
286            .push(values)
287    }
288
289    pub(super) fn solve_row(&self, slot: usize, row: usize) -> RuntimeResult<&[Complex64]> {
290        self.solve_row_slots
291            .get(slot)
292            .ok_or(RuntimeError::InvalidCache {
293                expected: self.solve_row_slots.len(),
294                actual: slot + 1,
295            })?
296            .row(row)
297    }
298}
299
300impl CpuPlan {
301    /// Materializes the event-dependent cache for a batch.
302    ///
303    /// # Errors
304    ///
305    /// Returns [`RuntimeError`] when required columns are missing, expression
306    /// shapes are invalid, cache construction fails, or a matrix is singular.
307    ///
308    /// # Panics
309    ///
310    /// Panics if a node selected by the validated cache plan was not evaluated.
311    pub fn cache_event_batch(&self, batch: &EventBatch) -> RuntimeResult<CpuBatchCache> {
312        let event_columns = self.event_columns(batch.schema())?;
313        let mut cache = CpuBatchCache::new(
314            &self.cache_plan,
315            &self.factor_matrices,
316            &self.solve_row_keys,
317            batch.len(),
318        )?;
319        if self.scalar_cache_supported() {
320            self.cache_scalar_batch(batch, &event_columns, &mut cache)?;
321            cache.set_weights((0..batch.len()).map(|row| batch.weights_at(row)).collect());
322            return Ok(cache);
323        }
324        for row in 0..batch.len() {
325            let values = self.evaluate_cache_values_for_row(batch, row, &event_columns)?;
326            for (slot, entry) in self.cache_plan.entries().iter().enumerate() {
327                let value = values[entry.node().index()]
328                    .as_ref()
329                    .expect("cacheable node should have been evaluated")
330                    .clone();
331                cache.push(slot, value)?;
332            }
333            for plan in &self.solve_row_matrices {
334                let (rows, cols, values) = matrix_at_optional(&values, plan.matrix().index())?;
335                if rows != plan.dimension() || cols != plan.dimension() {
336                    return Err(RuntimeError::InvalidShape {
337                        index: plan.matrix().index(),
338                        message: format!(
339                            "specialized solve expected a {}x{} matrix, got {rows}x{cols}",
340                            plan.dimension(),
341                            plan.dimension()
342                        ),
343                    });
344                }
345                let transpose_factor = DMatrix::from_row_slice(rows, cols, values).transpose().lu();
346                for (slot, index) in plan.rows() {
347                    let mut basis = DVector::zeros(plan.dimension());
348                    basis[*index] = Complex64::ONE;
349                    let inverse_row = transpose_factor
350                        .solve(&basis)
351                        .ok_or(RuntimeError::SingularMatrix(plan.matrix().index()))?;
352                    cache.push_solve_row(*slot, inverse_row.iter().copied())?;
353                }
354            }
355            for (slot, (matrix, _)) in self.factor_matrices.iter().enumerate() {
356                let (rows, cols, values) = matrix_at_optional(&values, matrix.index())?;
357                cache.push_factor(slot, DMatrix::from_row_slice(rows, cols, values).lu())?;
358            }
359        }
360        cache.set_weights((0..batch.len()).map(|row| batch.weights_at(row)).collect());
361        Ok(cache)
362    }
363
364    fn scalar_cache_supported(&self) -> bool {
365        self.factor_matrices.is_empty()
366            && self.solve_row_matrices.is_empty()
367            && self
368                .cache_plan
369                .entries()
370                .iter()
371                .all(|entry| matches!(entry.value_kind(), ValueKind::Real | ValueKind::Complex))
372            && self.cache_materialization_nodes.iter().all(|id| {
373                matches!(
374                    self.graph.node(*id),
375                    Some(
376                        ExprNode::RealConst(_)
377                            | ExprNode::ComplexConst(_)
378                            | ExprNode::EventScalar(_)
379                            | ExprNode::EventP4Component { .. }
380                            | ExprNode::Unary { .. }
381                            | ExprNode::Binary { .. }
382                            | ExprNode::NaryAdd { .. }
383                            | ExprNode::NaryMul { .. }
384                            | ExprNode::Complex { .. }
385                    )
386                )
387            })
388    }
389
390    fn cache_scalar_batch(
391        &self,
392        batch: &EventBatch,
393        event_columns: &[Option<EventColumn>],
394        cache: &mut CpuBatchCache,
395    ) -> RuntimeResult<()> {
396        let mut values = vec![Complex64::ZERO; self.graph.nodes().len()];
397        for row in 0..batch.len() {
398            for id in &self.cache_materialization_nodes {
399                let index = id.index();
400                values[index] = match &self.graph.nodes()[index] {
401                    ExprNode::RealConst(value) => Complex64::from(*value),
402                    ExprNode::ComplexConst(value) => *value,
403                    ExprNode::EventScalar(name) => {
404                        let Some(EventColumn::Scalar(col)) = event_columns[index] else {
405                            return Err(RuntimeError::MissingEventColumn(name.to_string()));
406                        };
407                        Complex64::from(batch.scalar_at(col, row))
408                    }
409                    ExprNode::EventP4Component { name, component } => {
410                        let Some(EventColumn::P4Component { col, .. }) = event_columns[index]
411                        else {
412                            return Err(RuntimeError::MissingEventColumn(name.to_string()));
413                        };
414                        let p4 = batch.p4_at(col, row);
415                        Complex64::from(match component {
416                            P4Component::Px => p4.px,
417                            P4Component::Py => p4.py,
418                            P4Component::Pz => p4.pz,
419                            P4Component::E => p4.e,
420                        })
421                    }
422                    ExprNode::Unary { op, input } => eval_unary(*op, values[input.index()]),
423                    ExprNode::Binary { op, lhs, rhs } => {
424                        eval_binary(*op, values[lhs.index()], values[rhs.index()])
425                    }
426                    ExprNode::NaryAdd { terms } => {
427                        terms.iter().map(|term| values[term.index()]).sum()
428                    }
429                    ExprNode::NaryMul { factors } => factors
430                        .iter()
431                        .map(|factor| values[factor.index()])
432                        .product(),
433                    ExprNode::Complex { re, im } => {
434                        Complex64::new(values[re.index()].re, values[im.index()].re)
435                    }
436                    _ => unreachable!("scalar cache support was checked"),
437                };
438            }
439            for (slot, entry) in self.cache_plan.entries().iter().enumerate() {
440                cache.push(slot, Value::Scalar(values[entry.node().index()]))?;
441            }
442        }
443        Ok(())
444    }
445}
446
447/// A cached event batch and its associated weights.
448#[derive(Clone, Debug)]
449pub struct CpuCachedBatch {
450    pub(super) cache: CpuBatchCache,
451}
452
453impl CpuCachedBatch {
454    pub(crate) fn from_cache(cache: CpuBatchCache) -> Self {
455        Self { cache }
456    }
457
458    /// Returns the underlying materialized cache.
459    pub fn cache(&self) -> &CpuBatchCache {
460        &self.cache
461    }
462
463    /// Returns the number of events.
464    pub fn len(&self) -> usize {
465        self.cache.len()
466    }
467
468    /// Returns whether the batch contains no events.
469    pub fn is_empty(&self) -> bool {
470        self.cache.is_empty()
471    }
472
473    /// Returns per-event weights.
474    pub fn weights(&self) -> &[f64] {
475        self.cache.weights()
476    }
477
478    /// Returns the sum of event weights.
479    pub fn sum_weights(&self) -> f64 {
480        self.cache.sum_weights()
481    }
482
483    /// Estimates retained heap memory, in bytes.
484    pub fn resident_bytes(&self) -> usize {
485        self.cache.resident_bytes()
486    }
487}
488
489/// A dataset whose event-dependent model values are fully cached in memory.
490#[derive(Clone, Debug, Default)]
491pub struct CpuCachedDataset {
492    pub(super) batches: Vec<CpuCachedBatch>,
493    pub(super) sum_weights: f64,
494}
495
496impl PreparedDatasetStats {
497    pub(crate) fn new(
498        local_events: usize,
499        global_events: usize,
500        local_batches: usize,
501        sum_weights: f64,
502        resident_bytes: usize,
503        storage: CacheStorage,
504    ) -> Self {
505        Self {
506            local_events,
507            global_events,
508            local_batches,
509            sum_weights,
510            resident_bytes,
511            storage,
512        }
513    }
514
515    /// Returns the number of events assigned to this rank.
516    pub fn local_events(&self) -> usize {
517        self.local_events
518    }
519
520    /// Returns the total number of events across all ranks.
521    pub fn global_events(&self) -> usize {
522        self.global_events
523    }
524
525    /// Returns the number of batches assigned to this rank.
526    pub fn local_batches(&self) -> usize {
527        self.local_batches
528    }
529
530    /// Returns the total event-weight sum across all ranks.
531    pub fn sum_weights(&self) -> f64 {
532        self.sum_weights
533    }
534
535    /// Returns the number of bytes retained for prepared data on this rank.
536    pub fn resident_bytes(&self) -> usize {
537        self.resident_bytes
538    }
539
540    /// Returns the dataset's cache-storage policy.
541    pub fn storage(&self) -> CacheStorage {
542        self.storage
543    }
544}
545
546#[derive(Clone)]
547/// A dataset prepared according to its [`CacheStorage`] policy.
548///
549/// Resident datasets own all event-dependent cache values. Streaming datasets retain the source
550/// and read plan and rebuild transient batch caches on every reduction.
551pub enum CpuPreparedDataset {
552    /// A dataset whose event caches are resident in memory.
553    Resident {
554        /// Fully cached dataset.
555        dataset: Arc<CpuCachedDataset>,
556        /// Preparation statistics.
557        stats: PreparedDatasetStats,
558        /// Persistent host-memory reservation shared by clones.
559        memory_lease: MemoryLease,
560    },
561    /// A dataset whose event caches are rebuilt while streaming.
562    Streaming {
563        /// Source dataset.
564        dataset: Dataset,
565        /// Read plan used for each pass.
566        read_plan: ReadPlan,
567        /// Preparation statistics.
568        stats: PreparedDatasetStats,
569        /// Peak transient bytes reserved during each reduction.
570        transient_bytes: u64,
571    },
572}
573
574impl std::fmt::Debug for CpuPreparedDataset {
575    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
576        formatter
577            .debug_struct("CpuPreparedDataset")
578            .field("stats", self.stats())
579            .finish_non_exhaustive()
580    }
581}
582
583impl CpuPreparedDataset {
584    /// Returns statistics collected while preparing the dataset.
585    pub fn stats(&self) -> &PreparedDatasetStats {
586        match self {
587            Self::Resident { stats, .. } | Self::Streaming { stats, .. } => stats,
588        }
589    }
590}
591
592impl CpuCachedDataset {
593    pub(crate) fn from_parts(batches: Vec<CpuCachedBatch>, sum_weights: f64) -> Self {
594        Self {
595            batches,
596            sum_weights,
597        }
598    }
599
600    /// Returns the cached batches.
601    pub fn batches(&self) -> &[CpuCachedBatch] {
602        &self.batches
603    }
604
605    /// Returns the total number of cached events.
606    pub fn len(&self) -> usize {
607        self.batches.iter().map(CpuCachedBatch::len).sum()
608    }
609
610    /// Returns whether the dataset contains no events.
611    pub fn is_empty(&self) -> bool {
612        self.batches.iter().all(CpuCachedBatch::is_empty)
613    }
614
615    /// Returns the sum of all event weights.
616    pub fn sum_weights(&self) -> f64 {
617        self.sum_weights
618    }
619
620    /// Estimates retained heap memory, in bytes.
621    pub fn resident_bytes(&self) -> usize {
622        self.batches
623            .iter()
624            .map(CpuCachedBatch::resident_bytes)
625            .sum()
626    }
627}
628
629#[derive(Clone, Debug)]
630pub(super) struct CachedFactorSlot {
631    dimension: usize,
632    factors: Vec<DynamicLu>,
633}
634
635#[derive(Clone, Debug)]
636pub(crate) struct CachedSolveRowSlot {
637    #[cfg_attr(not(feature = "jit"), allow(dead_code))]
638    pub(crate) dimension: usize,
639    pub(crate) values: FlatRows<Complex64>,
640}
641
642impl CachedSolveRowSlot {
643    fn new(dimension: usize, events: usize) -> RuntimeResult<Self> {
644        Ok(Self {
645            dimension,
646            values: FlatRows::try_with_capacity(dimension, events)?,
647        })
648    }
649
650    fn push(&mut self, values: impl IntoIterator<Item = Complex64>) -> RuntimeResult<()> {
651        self.values.push_row(values)
652    }
653
654    fn row(&self, row: usize) -> RuntimeResult<&[Complex64]> {
655        self.values.row(row)
656    }
657
658    fn resident_bytes(&self) -> usize {
659        self.values.capacity() * size_of::<Complex64>()
660    }
661}
662
663impl CachedFactorSlot {
664    fn new(dimension: usize) -> Self {
665        Self {
666            dimension,
667            factors: Vec::new(),
668        }
669    }
670
671    fn push(&mut self, factor: DynamicLu) -> RuntimeResult<()> {
672        self.factors.push(factor);
673        Ok(())
674    }
675
676    fn factor(&self, row: usize) -> RuntimeResult<&DynamicLu> {
677        self.factors
678            .get(row)
679            .ok_or_else(|| RuntimeError::InvalidShape {
680                index: row,
681                message: format!(
682                    "factor row {row} out of bounds for len {}",
683                    self.factors.len()
684                ),
685            })
686    }
687
688    fn resident_bytes(&self) -> usize {
689        self.factors.capacity()
690            * (self.dimension * self.dimension * size_of::<Complex64>()
691                + self.dimension * size_of::<usize>())
692    }
693}
694
695#[derive(Clone, Debug, PartialEq)]
696pub(crate) enum CachedSlot {
697    Real(Vec<f64>),
698    Complex(Vec<Complex64>),
699    Vector {
700        len: usize,
701        values: FlatRows<Complex64>,
702    },
703    Matrix {
704        rows: usize,
705        cols: usize,
706        values: FlatRows<Complex64>,
707    },
708}
709
710impl CachedSlot {
711    #[cfg(feature = "jit")]
712    pub(crate) fn values_ptr(&self) -> *const u8 {
713        match self {
714            Self::Real(values) => values.as_ptr().cast(),
715            Self::Complex(values) => values.as_ptr().cast(),
716            Self::Vector { values, .. } | Self::Matrix { values, .. } => values.as_ptr().cast(),
717        }
718    }
719
720    #[cfg(feature = "jit")]
721    pub(crate) fn width(&self) -> usize {
722        match self {
723            Self::Real(_) | Self::Complex(_) => 1,
724            Self::Vector { values, .. } | Self::Matrix { values, .. } => values.width(),
725        }
726    }
727
728    fn new(kind: ValueKind, events: usize) -> RuntimeResult<Self> {
729        Ok(match kind {
730            ValueKind::Real => Self::Real(Vec::with_capacity(events)),
731            ValueKind::Complex => Self::Complex(Vec::with_capacity(events)),
732            ValueKind::Vector { len } => Self::Vector {
733                len,
734                values: FlatRows::try_with_capacity(len, events)?,
735            },
736            ValueKind::Matrix { rows, cols } => Self::Matrix {
737                rows,
738                cols,
739                values: FlatRows::try_with_capacity(
740                    rows.checked_mul(cols)
741                        .ok_or_else(|| RuntimeError::InvalidShape {
742                            index: rows,
743                            message: format!("matrix width overflowed for {rows}x{cols}"),
744                        })?,
745                    events,
746                )?,
747            },
748        })
749    }
750
751    pub(crate) fn resident_bytes(&self) -> usize {
752        match self {
753            Self::Real(values) => values.capacity() * size_of::<f64>(),
754            Self::Complex(values) => values.capacity() * size_of::<Complex64>(),
755            Self::Vector { values, .. } | Self::Matrix { values, .. } => {
756                values.capacity() * size_of::<Complex64>()
757            }
758        }
759    }
760
761    fn push(&mut self, value: Value) -> RuntimeResult<()> {
762        match (self, value) {
763            (Self::Real(values), Value::Scalar(value)) => {
764                values.push(value.re);
765                Ok(())
766            }
767            (Self::Complex(values), Value::Scalar(value)) => {
768                values.push(value);
769                Ok(())
770            }
771            (Self::Vector { len, values }, Value::Vector(value)) if *len == value.len() => {
772                values.push_row(value)
773            }
774            (
775                Self::Matrix { rows, cols, values },
776                Value::Matrix {
777                    rows: value_rows,
778                    cols: value_cols,
779                    values: value,
780                },
781            ) if *rows == value_rows && *cols == value_cols => values.push_row(value),
782            (_, value) => Err(RuntimeError::InvalidShape {
783                index: 0,
784                message: format!("cached value kind did not match slot: {}", value.kind()),
785            }),
786        }
787    }
788
789    pub(super) fn value(&self, row: usize) -> RuntimeResult<Value> {
790        match self {
791            Self::Real(values) => values
792                .get(row)
793                .copied()
794                .map(Complex64::from)
795                .map(Value::Scalar)
796                .ok_or_else(|| RuntimeError::InvalidShape {
797                    index: row,
798                    message: format!("cache row {row} out of bounds"),
799                }),
800            Self::Complex(values) => values.get(row).copied().map(Value::Scalar).ok_or_else(|| {
801                RuntimeError::InvalidShape {
802                    index: row,
803                    message: format!("cache row {row} out of bounds"),
804                }
805            }),
806            Self::Vector { values, .. } => {
807                values.row(row).map(|value| Value::Vector(value.to_vec()))
808            }
809            Self::Matrix { rows, cols, values } => values.row(row).map(|value| Value::Matrix {
810                rows: *rows,
811                cols: *cols,
812                values: value.to_vec(),
813            }),
814        }
815    }
816
817    fn scalar(&self, row: usize) -> RuntimeResult<Complex64> {
818        match self {
819            Self::Real(values) => values
820                .get(row)
821                .copied()
822                .map(Complex64::from)
823                .ok_or_else(|| RuntimeError::InvalidShape {
824                    index: row,
825                    message: format!("cache row {row} out of bounds"),
826                }),
827            Self::Complex(values) => {
828                values
829                    .get(row)
830                    .copied()
831                    .ok_or_else(|| RuntimeError::InvalidShape {
832                        index: row,
833                        message: format!("cache row {row} out of bounds"),
834                    })
835            }
836            Self::Vector { .. } | Self::Matrix { .. } => Err(RuntimeError::TypeMismatch {
837                index: row,
838                expected: "scalar",
839                actual: match self {
840                    Self::Vector { .. } => "vector",
841                    Self::Matrix { .. } => "matrix",
842                    Self::Real(_) | Self::Complex(_) => unreachable!(),
843                },
844            }),
845        }
846    }
847
848    fn real_range(&self, start: usize, end: usize) -> RuntimeResult<&[f64]> {
849        match self {
850            Self::Real(values) => {
851                values
852                    .get(start..end)
853                    .ok_or_else(|| RuntimeError::InvalidShape {
854                        index: start,
855                        message: format!("cache range {start}..{end} out of bounds"),
856                    })
857            }
858            Self::Complex(_) | Self::Vector { .. } | Self::Matrix { .. } => {
859                Err(RuntimeError::TypeMismatch {
860                    index: start,
861                    expected: "real scalar",
862                    actual: match self {
863                        Self::Complex(_) => "complex scalar",
864                        Self::Vector { .. } => "vector",
865                        Self::Matrix { .. } => "matrix",
866                        Self::Real(_) => unreachable!(),
867                    },
868                })
869            }
870        }
871    }
872
873    fn complex_range(&self, start: usize, end: usize) -> RuntimeResult<&[Complex64]> {
874        match self {
875            Self::Complex(values) => {
876                values
877                    .get(start..end)
878                    .ok_or_else(|| RuntimeError::InvalidShape {
879                        index: start,
880                        message: format!("cache range {start}..{end} out of bounds"),
881                    })
882            }
883            Self::Real(_) | Self::Vector { .. } | Self::Matrix { .. } => {
884                Err(RuntimeError::TypeMismatch {
885                    index: start,
886                    expected: "complex scalar",
887                    actual: match self {
888                        Self::Real(_) => "real scalar",
889                        Self::Vector { .. } => "vector",
890                        Self::Matrix { .. } => "matrix",
891                        Self::Complex(_) => unreachable!(),
892                    },
893                })
894            }
895        }
896    }
897}