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#[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#[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 pub fn len(&self) -> usize {
87 self.len
88 }
89
90 pub fn is_empty(&self) -> bool {
92 self.len == 0
93 }
94
95 pub fn weights(&self) -> &[f64] {
97 &self.weights
98 }
99
100 pub fn sum_weights(&self) -> f64 {
102 self.sum_weights
103 }
104
105 #[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 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 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#[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 pub fn cache(&self) -> &CpuBatchCache {
460 &self.cache
461 }
462
463 pub fn len(&self) -> usize {
465 self.cache.len()
466 }
467
468 pub fn is_empty(&self) -> bool {
470 self.cache.is_empty()
471 }
472
473 pub fn weights(&self) -> &[f64] {
475 self.cache.weights()
476 }
477
478 pub fn sum_weights(&self) -> f64 {
480 self.cache.sum_weights()
481 }
482
483 pub fn resident_bytes(&self) -> usize {
485 self.cache.resident_bytes()
486 }
487}
488
489#[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 pub fn local_events(&self) -> usize {
517 self.local_events
518 }
519
520 pub fn global_events(&self) -> usize {
522 self.global_events
523 }
524
525 pub fn local_batches(&self) -> usize {
527 self.local_batches
528 }
529
530 pub fn sum_weights(&self) -> f64 {
532 self.sum_weights
533 }
534
535 pub fn resident_bytes(&self) -> usize {
537 self.resident_bytes
538 }
539
540 pub fn storage(&self) -> CacheStorage {
542 self.storage
543 }
544}
545
546#[derive(Clone)]
547pub enum CpuPreparedDataset {
552 Resident {
554 dataset: Arc<CpuCachedDataset>,
556 stats: PreparedDatasetStats,
558 memory_lease: MemoryLease,
560 },
561 Streaming {
563 dataset: Dataset,
565 read_plan: ReadPlan,
567 stats: PreparedDatasetStats,
569 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 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 pub fn batches(&self) -> &[CpuCachedBatch] {
602 &self.batches
603 }
604
605 pub fn len(&self) -> usize {
607 self.batches.iter().map(CpuCachedBatch::len).sum()
608 }
609
610 pub fn is_empty(&self) -> bool {
612 self.batches.iter().all(CpuCachedBatch::is_empty)
613 }
614
615 pub fn sum_weights(&self) -> f64 {
617 self.sum_weights
618 }
619
620 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}