Skip to main content

only_brain/
bvector.rs

1use std::ops;
2
3use nalgebra::{SimdRealField, SVector};
4
5/// This module provides a vector type `BVector` that is optimized for compile-time
6/// dimension and uses SIMD operations for performance.
7#[derive(Debug, Clone, PartialEq)]
8pub struct BVector<T: SimdRealField + Copy, const N: usize> {
9    pub data_vector: SVector<T, N>,
10}
11
12impl<T: SimdRealField + Copy, const N: usize> BVector<T, N> {
13    /// Constructs a vector filled with `element` of a dimension `N`.
14    pub fn from_element(element: T) -> Self {
15        BVector {
16            data_vector: SVector::<T, N>::from_element(element),
17        }
18    }
19
20    /// Constructs the vector from fixed-size array.
21    pub fn from_array(data: [T; N]) -> Self
22    where
23        T: Copy,
24    {
25        BVector {
26            data_vector: SVector::<T, N>::from_row_slice(&data),
27        }
28    }
29
30    pub fn dot(&self, other: &Self) -> T {
31        self.data_vector.dot(&other.data_vector)
32    }
33
34    pub fn len(&self) -> usize {
35        N
36    }
37
38    /// Whether the vector has no elements, which is only the case when `N` is 0.
39    pub fn is_empty(&self) -> bool {
40        N == 0
41    }
42
43    pub fn get(&self, index: usize) -> T {
44        assert!(index < N, "Index out of bounds");
45        self.data_vector[index]
46    }
47}
48
49impl<T: SimdRealField + Copy, const N: usize> ops::Add for BVector<T, N> {
50    type Output = BVector<T, N>;
51
52    fn add(self, rhs: Self) -> Self::Output {
53        BVector {
54            data_vector: self.data_vector + rhs.data_vector,
55        }
56    }
57}
58
59/// Macro to construct a `BVector` with compile-time dimension inferred from the
60/// number of elements provided.
61///
62/// Example:
63/// - `let v = bvector![1.0, 2.0, 3.0]; // BVector<f64, 3>`
64/// - `let v = bvector![2.0, -1.5]; // BVector<f64, 2>`
65#[macro_export]
66macro_rules! bvector {
67    ($($x:expr),+ $(,)?) => {
68        {
69            // Let the compiler infer both element type `T` and length `N`.
70            $crate::BVector::from_array([$( $x ),+])
71        }
72    };
73}
74
75
76#[cfg(test)]
77mod tests {
78    use super::*;
79
80    #[test]
81    fn from_element_fills_every_position() {
82        let v = BVector::<f64, 3>::from_element(1.5);
83
84        assert_eq!((v.get(0), v.get(1), v.get(2)), (1.5, 1.5, 1.5));
85    }
86
87    #[test]
88    fn the_macro_keeps_elements_in_order() {
89        let v = bvector![1.0, 2.0, 3.0];
90
91        assert_eq!(v, BVector::from_array([1.0, 2.0, 3.0]));
92        assert_eq!((v.get(0), v.get(1), v.get(2)), (1.0, 2.0, 3.0));
93    }
94
95    #[test]
96    fn the_macro_accepts_a_trailing_comma() {
97        assert_eq!(bvector![1.0, 2.0,], bvector![1.0, 2.0]);
98    }
99
100    #[test]
101    fn len_is_the_compile_time_dimension() {
102        assert_eq!(bvector![1.0, 2.0, 3.0].len(), 3);
103        assert!(!bvector![1.0].is_empty());
104        assert!(BVector::<f64, 0>::from_array([]).is_empty());
105    }
106
107    #[test]
108    fn dot_sums_the_element_wise_products() {
109        assert_eq!(bvector![1.0, 2.0, 3.0].dot(&bvector![4.0, -5.0, 6.0]), 12.0);
110    }
111
112    #[test]
113    fn add_is_element_wise() {
114        assert_eq!(bvector![1.0, 2.0] + bvector![0.5, -3.0], bvector![1.5, -1.0]);
115    }
116
117    #[test]
118    #[should_panic(expected = "Index out of bounds")]
119    fn get_rejects_an_index_past_the_end() {
120        bvector![1.0, 2.0].get(2);
121    }
122}