Skip to main content

core_utils/circuit/latest/
slice.rs

1use itertools::izip;
2use serde::{Deserialize, Serialize};
3
4use crate::circuit::errors::SliceError;
5
6/// A general slicing structure which can represent a single index, a strided 1d range, a strided 2d
7/// range or a vector of slices.
8#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
9pub struct Slice(SliceEnum);
10
11#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
12#[repr(C)]
13enum SliceEnum {
14    /// A single index slice
15    Single(u32),
16    /// A slice with indices given by
17    /// ```text
18    /// (0..size).map(|i| start + i * step)
19    /// ```
20    Range { start: u32, size: u32, step: i64 },
21    /// A slice with indices given by
22    /// ```text
23    /// (0..size1).flat_map(|i| (0..size2).map(|j| start + step1 * i + step2 * j))
24    /// ```
25    Range2d {
26        start: u32,
27        size1: u32,
28        step1: i64,
29        size2: u32,
30        step2: i64,
31    },
32    /// A slice with indices given by a vector of slices.
33    RangeVec(Vec<SliceEnum>),
34}
35
36impl Slice {
37    pub fn empty() -> Self {
38        Self(SliceEnum::RangeVec(vec![]))
39    }
40
41    pub fn single(index: u32) -> Self {
42        Self(SliceEnum::Single(index))
43    }
44
45    pub fn range(start: u32, size: u32, step: i64) -> Result<Self, SliceError> {
46        validate_range_bounds(start, size, step)?;
47        Ok(Self(SliceEnum::Range { start, size, step }))
48    }
49
50    pub fn shift_start(&mut self, delta: u32) {
51        self.0.shift_start(delta);
52    }
53
54    pub fn range2d(
55        start: u32,
56        size1: u32,
57        size2: u32,
58        step1: i64,
59        step2: i64,
60    ) -> Result<Self, SliceError> {
61        validate_range_2d_bounds(start, size1, step1, size2, step2)?;
62        Ok(Self(SliceEnum::Range2d {
63            start,
64            size1,
65            step1,
66            size2,
67            step2,
68        }))
69    }
70
71    pub fn append(&mut self, other: Self) {
72        match (&mut self.0, other.0) {
73            (SliceEnum::RangeVec(v), SliceEnum::RangeVec(v1)) => v.extend(v1),
74            (SliceEnum::RangeVec(v), slice) => v.push(slice),
75            (slice, SliceEnum::RangeVec(mut v1)) => {
76                v1.insert(0, slice.clone());
77                *slice = SliceEnum::RangeVec(v1);
78            }
79            (slice, slice1) => *slice = SliceEnum::RangeVec(vec![slice.clone(), slice1]),
80        }
81    }
82
83    pub fn get_indices(&self) -> Vec<u32> {
84        self.0.get_indices()
85    }
86
87    pub fn is_empty(&self) -> bool {
88        self.len() == 0
89    }
90
91    pub fn len(&self) -> u32 {
92        self.0.len()
93    }
94
95    pub fn from_indices(indices: Vec<u32>) -> Self {
96        Self(SliceEnum::from_indices(indices))
97    }
98
99    pub fn optimize(self) -> Self {
100        Self::from_indices(self.get_indices())
101    }
102}
103
104fn validate_bounds(min_index: i128, max_index: i128) -> Result<(), SliceError> {
105    if min_index < 0 {
106        return Err(SliceError::NegativeIndex(min_index));
107    }
108    if max_index > i128::from(u32::MAX) {
109        return Err(SliceError::IndexOutOfBounds {
110            found: max_index,
111            max: u32::MAX,
112        });
113    }
114    Ok(())
115}
116
117#[inline]
118fn range_index(start: u32, step: i64, i: i64) -> i128 {
119    i128::from(start) + i128::from(step) * i128::from(i)
120}
121
122#[inline]
123fn range_2d_index(start: u32, step1: i64, i: i64, step2: i64, j: i64) -> i128 {
124    i128::from(start) + i128::from(step1) * i128::from(i) + i128::from(step2) * i128::from(j)
125}
126
127fn validate_range_bounds(start: u32, size: u32, step: i64) -> Result<(), SliceError> {
128    if size == 0 {
129        return Ok(());
130    }
131    let last = i64::from(size - 1);
132    let first = i128::from(start);
133    let end = range_index(start, step, last);
134    validate_bounds(first.min(end), first.max(end))
135}
136
137fn validate_range_2d_bounds(
138    start: u32,
139    size1: u32,
140    step1: i64,
141    size2: u32,
142    step2: i64,
143) -> Result<(), SliceError> {
144    if size1 == 0 || size2 == 0 {
145        return Ok(());
146    }
147    let i_last = i64::from(size1 - 1);
148    let j_last = i64::from(size2 - 1);
149
150    let corners = [
151        range_2d_index(start, step1, 0, step2, 0),
152        range_2d_index(start, step1, i_last, step2, 0),
153        range_2d_index(start, step1, 0, step2, j_last),
154        range_2d_index(start, step1, i_last, step2, j_last),
155    ];
156
157    let min_index = corners.into_iter().min().unwrap_or(0);
158    let max_index = corners.into_iter().max().unwrap_or(0);
159    validate_bounds(min_index, max_index)
160}
161
162#[inline]
163fn to_u32_index(index: i128) -> u32 {
164    u32::try_from(index).unwrap_or_else(|_| panic!("slice index out of bounds: {index}"))
165}
166
167fn generate_range_indices(start: u32, size: u32, step: i64) -> impl Iterator<Item = u32> {
168    (0..i64::from(size)).map(move |i| to_u32_index(range_index(start, step, i)))
169}
170
171fn generate_range_2d_indices(
172    start: u32,
173    size1: u32,
174    step1: i64,
175    size2: u32,
176    step2: i64,
177) -> impl Iterator<Item = u32> {
178    (0..i64::from(size1)).flat_map(move |i| {
179        (0..i64::from(size2)).map(move |j| to_u32_index(range_2d_index(start, step1, i, step2, j)))
180    })
181}
182
183impl SliceEnum {
184    fn get_indices(&self) -> Vec<u32> {
185        match self {
186            SliceEnum::Single(idx) => vec![*idx],
187            SliceEnum::Range { start, size, step } => {
188                generate_range_indices(*start, *size, *step).collect()
189            }
190            SliceEnum::Range2d {
191                start,
192                size1,
193                size2,
194                step1,
195                step2,
196            } => generate_range_2d_indices(*start, *size1, *step1, *size2, *step2).collect(),
197            SliceEnum::RangeVec(v) => v.iter().flat_map(|r| r.get_indices()).collect(),
198        }
199    }
200
201    pub fn len(&self) -> u32 {
202        match self {
203            SliceEnum::Single(_) => 1,
204            SliceEnum::Range { size, .. } => *size,
205            SliceEnum::Range2d { size1, size2, .. } => size1
206                .checked_mul(*size2)
207                .expect("slice length overflow for range2d"),
208            SliceEnum::RangeVec(v) => v.iter().fold(0u32, |acc, r| {
209                acc.checked_add(r.len())
210                    .expect("slice length overflow for range vector")
211            }),
212        }
213    }
214
215    /// Given a start index and a vector of deltas tries to find a slice (`Single`, `Range` or
216    /// `Range2d`) which generates the longest sequence `[index0, index0 + deltas[0], index0 +
217    /// deltas[0] + deltas[1], ...]`
218    fn match_largest_slice(start: u32, deltas: &[i64]) -> Self {
219        if deltas.is_empty() {
220            return Self::Single(start);
221        }
222
223        // The longest sequence of equal deltas generates a 1d range slice
224        // A 1d slice verifies: `deltas[..] = deltas[0] | deltas[0] | .. | deltas[0]`
225        let step_j = deltas[0];
226        let n_j = deltas.iter().skip(1).take_while(|&&d| d == step_j).count() + 2;
227
228        let mut res_slice = Self::Range {
229            start,
230            size: n_j as u32,
231            step: step_j,
232        };
233
234        if n_j < deltas.len() + 1 {
235            // If the sequence of deltas is not finished, try to match a 2d slice.
236            // A 2d slice verifies:
237            //  `deltas[..] = deltas[0..n_j] | deltas[0..n_j] | .. | deltas[0..n_j - 1]`
238            let exp_chunk = &deltas[0..n_j];
239            let chunks = deltas.chunks(n_j).skip(1);
240            let mut n_i = chunks
241                .take_while(|chunk| {
242                    izip!(exp_chunk, *chunk).take_while(|(e, d)| e == d).count() == n_j
243                })
244                .count()
245                + 1;
246            if let Some(chunk) = deltas.chunks(n_j).nth(n_i) {
247                if izip!(exp_chunk, chunk).take_while(|(e, d)| e == d).count() == n_j - 1 {
248                    n_i += 1;
249                }
250            }
251
252            if n_i > 1 {
253                let step_i = exp_chunk.iter().sum::<i64>();
254                res_slice = Self::Range2d {
255                    start,
256                    size1: n_i as u32,
257                    size2: n_j as u32,
258                    step1: step_i,
259                    step2: step_j,
260                };
261            }
262        }
263
264        res_slice
265    }
266
267    /// Reduces the current slice to a slice with at most `new_size` indices.
268    fn reduce(&mut self, max_size: u32) {
269        assert!(max_size > 0);
270        match self {
271            SliceEnum::Single(_) => {}
272            SliceEnum::Range { start, size, .. } => {
273                if max_size < *size {
274                    if max_size == 1 {
275                        *self = SliceEnum::Single(*start);
276                    } else {
277                        *size = max_size;
278                    }
279                }
280            }
281            SliceEnum::Range2d {
282                start,
283                size1,
284                size2,
285                step2,
286                ..
287            } => {
288                if max_size < *size1 * *size2 {
289                    if max_size == 1 {
290                        *self = SliceEnum::Single(*start);
291                    } else if max_size <= *size2 {
292                        *self = SliceEnum::Range {
293                            start: *start,
294                            size: max_size,
295                            step: *step2,
296                        }
297                    } else if max_size / *size2 == 1 {
298                        *self = SliceEnum::Range {
299                            start: *start,
300                            size: *size2,
301                            step: *step2,
302                        }
303                    } else {
304                        *size1 = max_size / *size2;
305                    }
306                }
307            }
308            SliceEnum::RangeVec(_) => {}
309        }
310    }
311
312    fn match_slices(mut max_len_slices: Vec<Self>) -> Vec<Self> {
313        let mut res = vec![]; // result slices with absolute start indices
314        let mut ranges_to_visit = vec![(0, max_len_slices.len())]; // start with full range
315        while let Some((start, end)) = ranges_to_visit.pop() {
316            // Find the slice which generates the longest sequence of indices in the current range
317            // `[start, end)`
318            let (slice_pos, slice) = max_len_slices[start..end]
319                .iter()
320                .enumerate()
321                .max_by_key(|(pos, slice)| (slice.len(), end - pos)) // `end - pos` is used to return the first maximum
322                .unwrap();
323            let slice_start = start + slice_pos; // to absolute position
324            let slice_end = slice_start + slice.len() as usize;
325
326            // Store the max slice for the result
327            res.push((slice_start, slice.clone()));
328
329            // Add left and right ranges to visit if they are not empty
330            if start < slice_start {
331                // Reduce the length of the slices on before the max slice to not overlap with the
332                // max slice
333                max_len_slices[start..slice_start]
334                    .iter_mut()
335                    .enumerate()
336                    .for_each(|(pos, slice)| slice.reduce((slice_pos - pos) as u32));
337
338                ranges_to_visit.push((start, slice_start));
339            }
340            if slice_end < end {
341                ranges_to_visit.push((slice_end, end));
342            }
343        }
344
345        res.sort_by_key(|(start, _)| *start);
346        res.into_iter().map(|(_, slice)| slice).collect()
347    }
348
349    /// Given a vector of indices tries to find a minimal number of slices
350    /// (`Single`, `Range` or `Range2d`) which generates the same sequence of indices.
351    ///
352    /// The algorithm finds the largest slice in input sequence `indices` and then
353    /// recursively matches slices in the left-hand and right-hand size indices which are not
354    /// covered by the largest slice.
355    pub fn from_indices(indices: Vec<u32>) -> Self {
356        if indices.is_empty() {
357            return Self::RangeVec(vec![]);
358        }
359
360        let deltas = indices
361            .windows(2)
362            .map(|w| w[1] as i64 - w[0] as i64)
363            .collect::<Vec<_>>();
364        let max_slice_vec: Vec<_> = (0..indices.len())
365            .map(|i| Self::match_largest_slice(indices[i], &deltas[i..]))
366            .collect();
367
368        let optimized_slices = SliceEnum::match_slices(max_slice_vec);
369        if optimized_slices.len() == 1 {
370            optimized_slices[0].clone()
371        } else {
372            SliceEnum::RangeVec(optimized_slices)
373        }
374    }
375
376    pub fn shift_start(&mut self, delta: u32) {
377        match self {
378            SliceEnum::Single(idx) => {
379                *idx = idx
380                    .checked_add(delta)
381                    .expect("slice start overflow for single index");
382            }
383            SliceEnum::Range { start, .. } => {
384                *start = start
385                    .checked_add(delta)
386                    .expect("slice start overflow for range");
387            }
388            SliceEnum::Range2d { start, .. } => {
389                *start = start
390                    .checked_add(delta)
391                    .expect("slice start overflow for range2d");
392            }
393            SliceEnum::RangeVec(v) => v.iter_mut().for_each(|slice| slice.shift_start(delta)),
394        }
395    }
396}
397
398#[cfg(test)]
399mod tests {
400    use super::SliceEnum;
401    use crate::circuit::{errors::SliceError, Slice};
402
403    #[test]
404    fn test_slice_range() {
405        let range = SliceEnum::Range2d {
406            start: 0,
407            size1: 2,
408            size2: 3,
409            step1: 6,
410            step2: 1,
411        };
412        let expected = vec![0, 1, 2, 6, 7, 8];
413        assert_eq!(range.get_indices(), expected);
414
415        let range = SliceEnum::Range2d {
416            start: 0,
417            size1: 4,
418            size2: 2,
419            step1: 3,
420            step2: 1,
421        };
422        let expected = vec![0, 1, 3, 4, 6, 7, 9, 10];
423        assert_eq!(range.get_indices(), expected);
424
425        let range = SliceEnum::Range2d {
426            start: 0,
427            size1: 4,
428            size2: 2,
429            step1: 3,
430            step2: 2,
431        };
432        let expected = vec![0, 2, 3, 5, 6, 8, 9, 11];
433        assert_eq!(range.get_indices(), expected);
434
435        let range = SliceEnum::Range2d {
436            start: 2,
437            size1: 1,
438            size2: 4,
439            step1: 1,
440            step2: 3,
441        };
442        let expected = vec![2, 5, 8, 11];
443        assert_eq!(range.get_indices(), expected);
444    }
445
446    #[test]
447    fn test_slice_match_largest_slice() {
448        fn match_largest_slice(indices: &[u32]) -> SliceEnum {
449            SliceEnum::match_largest_slice(
450                indices[0],
451                &indices
452                    .windows(2)
453                    .map(|w| w[1] as i64 - w[0] as i64)
454                    .collect::<Vec<_>>(),
455            )
456        }
457
458        //// Full match
459        // single point slices
460        let indices = vec![0];
461        let slice = match_largest_slice(&indices);
462        assert_eq!(slice.get_indices(), indices);
463
464        let indices = vec![3];
465        let slice = match_largest_slice(&indices);
466        assert_eq!(slice.get_indices(), indices);
467
468        // 1d slices
469        let indices = vec![0, 1, 2, 3, 4];
470        let slice = match_largest_slice(&indices);
471        assert_eq!(slice.get_indices(), indices);
472
473        let indices = vec![5, 7, 9, 11, 13];
474        let slice = match_largest_slice(&indices);
475        assert_eq!(slice.get_indices(), indices);
476
477        let indices = vec![5, 6];
478        let slice = match_largest_slice(&indices);
479        assert_eq!(slice.get_indices(), indices);
480
481        let indices = vec![5, 2];
482        let slice = match_largest_slice(&indices);
483        assert_eq!(slice.get_indices(), indices[..2].to_vec());
484
485        // 2d slices
486        let indices = vec![0, 1, 2, 5, 6, 7, 10, 11, 12, 15, 16, 17]; // A[0..4][0..3] in a 4x5 matrix (row-major order)
487        let slice = match_largest_slice(&indices);
488        assert_eq!(slice.get_indices(), indices);
489
490        let indices = vec![2, 3, 4, 7, 8, 9]; // A[0..2][2..5] in a 2x5 matrix (row-major order)
491        let slice = match_largest_slice(&indices);
492        assert_eq!(slice.get_indices(), indices);
493
494        let indices = vec![0, 2, 8, 10]; // A[(0..3).step_by(2)][(0..4).step_by(2)] in a 3x4 matrix (row-major order)
495        let slice = match_largest_slice(&indices);
496        assert_eq!(slice.get_indices(), indices);
497
498        let indices = vec![10, 12, 5, 7, 0, 2]; // A[(0..3).reverse()][(0..3).step_by(2)] in a 3x5 matrix (row-major order)
499        let slice = match_largest_slice(&indices);
500        assert_eq!(slice.get_indices(), indices.to_vec());
501
502        //// Partial matches
503        // 1d slices
504        let indices = vec![0, 2, 4, 4, 5];
505        let slice = match_largest_slice(&indices);
506        assert_eq!(slice.get_indices(), indices[..3].to_vec());
507
508        // 2d slices
509        let indices = vec![0, 1, 3, 4, 5];
510        let slice = match_largest_slice(&indices);
511        assert_eq!(slice.get_indices(), indices[..4].to_vec());
512
513        let indices = vec![10, 12, 5, 7, 0, 2, 1];
514        let slice = match_largest_slice(&indices);
515        assert_eq!(slice.get_indices(), indices[..6].to_vec());
516
517        // Special cases
518        let indices = vec![1, 1, 0, 0, 1, 1, 0, 0];
519        let slice = match_largest_slice(&indices);
520        assert_eq!(slice.get_indices(), indices[..4].to_vec());
521    }
522
523    #[test]
524    fn test_slice_optimize() {
525        let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11];
526        let slice = Slice::from_indices(indices.clone());
527        assert_eq!(slice.get_indices(), indices);
528        assert_eq!(
529            slice.0,
530            SliceEnum::Range {
531                start: 0,
532                size: 12,
533                step: 1,
534            }
535        );
536
537        let indices = vec![19, 3, 4, 5, 6, 7, 8, 9, 10, 11];
538        let slice = Slice::from_indices(indices.clone());
539        assert_eq!(slice.get_indices(), indices);
540        assert_eq!(
541            slice.0,
542            SliceEnum::RangeVec(vec![
543                SliceEnum::Single(19),
544                SliceEnum::Range {
545                    start: 3,
546                    size: 9,
547                    step: 1
548                }
549            ])
550        );
551
552        let indices = vec![0, 1, 2, 19, 3, 4, 5, 6, 7, 8, 9, 10, 11];
553        let slice = Slice::from_indices(indices.clone());
554        assert_eq!(slice.get_indices(), indices);
555        assert_eq!(
556            slice.0,
557            SliceEnum::RangeVec(vec![
558                SliceEnum::Range {
559                    start: 0,
560                    size: 3,
561                    step: 1
562                },
563                SliceEnum::Single(19),
564                SliceEnum::Range {
565                    start: 3,
566                    size: 9,
567                    step: 1
568                }
569            ])
570        );
571
572        let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 19];
573        let slice = Slice::from_indices(indices.clone());
574        assert_eq!(slice.get_indices(), indices);
575        assert_eq!(
576            slice.0,
577            SliceEnum::RangeVec(vec![
578                SliceEnum::Range {
579                    start: 0,
580                    size: 10,
581                    step: 1
582                },
583                SliceEnum::Single(19),
584            ])
585        );
586
587        let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 19, 10, 11];
588        let slice = Slice::from_indices(indices.clone());
589        assert_eq!(slice.get_indices(), indices);
590        assert_eq!(
591            slice.0,
592            SliceEnum::RangeVec(vec![
593                SliceEnum::Range {
594                    start: 0,
595                    size: 10,
596                    step: 1
597                },
598                SliceEnum::Range {
599                    start: 19,
600                    size: 2,
601                    step: -9
602                },
603                SliceEnum::Single(11),
604            ])
605        );
606
607        // Large example, 4000 indices
608        let mut indices = Vec::new();
609        for _i in 0..1000 {
610            indices.extend(vec![0, 1, 1, 0]);
611        }
612        let slice = Slice::from_indices(indices.clone());
613        assert_eq!(slice.get_indices(), indices);
614    }
615
616    #[test]
617    fn test_slice_checked_range_bounds() {
618        assert_eq!(Slice::range(0, 2, -1), Err(SliceError::NegativeIndex(-1)));
619        assert_eq!(
620            Slice::range(u32::MAX, 2, 1),
621            Err(SliceError::IndexOutOfBounds {
622                found: i128::from(u32::MAX) + 1,
623                max: u32::MAX
624            })
625        );
626
627        let slice = Slice::range(u32::MAX - 1, 2, 1).unwrap();
628        assert_eq!(slice.get_indices(), vec![u32::MAX - 1, u32::MAX]);
629    }
630
631    #[test]
632    fn test_slice_checked_range2d_bounds() {
633        assert_eq!(
634            Slice::range2d(0, 2, 2, -1, 0),
635            Err(SliceError::NegativeIndex(-1))
636        );
637        assert_eq!(
638            Slice::range2d(u32::MAX, 2, 1, 1, 0),
639            Err(SliceError::IndexOutOfBounds {
640                found: i128::from(u32::MAX) + 1,
641                max: u32::MAX
642            })
643        );
644
645        let slice = Slice::range2d(u32::MAX - 1, 1, 2, 1, 1).unwrap();
646        assert_eq!(slice.get_indices(), vec![u32::MAX - 1, u32::MAX]);
647    }
648}