Skip to main content

draco_nd_vector/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use std::env;
4use syn::{parse_macro_input, LitInt};
5
6#[proc_macro]
7pub fn impl_ndvector_ops_for_dim(input: TokenStream) -> TokenStream {
8    let n = parse_macro_input!(input as LitInt)
9        .base10_parse::<usize>()
10        .unwrap();
11
12    let indices_add_assign = (0..n).map(|i| {
13        quote! { *self.data.get_unchecked_mut(#i) += *rhs.data.get_unchecked(#i); }
14    });
15    let indices_sub_assign = (0..n).map(|i| {
16        quote! { *self.data.get_unchecked_mut(#i) -= *rhs.data.get_unchecked(#i); }
17    });
18    let indices_mul_assign = (0..n).map(|i| {
19        quote! { *self.data.get_unchecked_mut(#i) *= rhs; }
20    });
21    let indices_div_assign = (0..n).map(|i| {
22        quote! { *self.data.get_unchecked_mut(#i) /= rhs; }
23    });
24    let indices_dot = (0..n).map(|i| {
25        quote! { result += *self.data.get_unchecked(#i) * *rhs.data.get_unchecked(#i); }
26    });
27    let indices_elem_mul = (0..n).map(|i| {
28        quote! { *self.data.get_unchecked_mut(#i) *= *rhs.data.get_unchecked(#i); }
29    });
30    let indices_elem_div = (0..n).map(|i| {
31        quote! { *self.data.get_unchecked_mut(#i) /= *rhs.data.get_unchecked(#i); }
32    });
33    let indices_partial_eq = (0..n).map(|i| {
34        quote! { result &= self.data.get_unchecked(#i).eq(rhs.data.get_unchecked(#i)); }
35    });
36    let indices_portable_to_bytes = (0..n).map(|i| {
37        quote! { result.extend((*self.data.get_unchecked(#i)).to_bytes()); }
38    });
39    let indices_portable_write_to = (0..n).map(|i| {
40        quote! { (*self.data.get_unchecked(#i)).write_to(writer); }
41    });
42    let indices_portable_read_from = (0..n).map(|i| {
43        quote! { *data.get_unchecked_mut(#i) = Data::read_from(reader)?; }
44    });
45
46    let expanded = quote! {
47        impl<T> std::ops::Add for NdVector<#n, T> 
48            where 
49                T: std::ops::Add<Output = T> + std::ops::AddAssign + Copy
50        {
51            type Output = NdVector<#n, T>;
52
53            fn add(mut self, rhs: Self) -> Self::Output {
54                self+=rhs;
55                self
56            }
57        }
58
59        impl<T> std::ops::AddAssign for NdVector<#n, T> 
60            where 
61                T: std::ops::AddAssign + Copy
62        {
63            fn add_assign(&mut self, rhs: Self) {
64                unsafe { #(#indices_add_assign)* }
65            }
66        }
67
68        impl<T> std::ops::Sub for NdVector<#n, T> 
69            where 
70                T: std::ops::Sub<Output = T> + std::ops::SubAssign + Copy
71        {
72            type Output = NdVector<#n, T>;
73
74            fn sub(mut self, rhs: Self) -> Self::Output {
75                self -= rhs;
76                self
77            }
78        }
79
80        impl<T> std::ops::SubAssign for NdVector<#n, T> 
81        where 
82            T: std::ops::SubAssign + Copy
83        {
84            fn sub_assign(&mut self, rhs: Self) {
85                unsafe { #(#indices_sub_assign)* }
86            }
87        }
88
89        impl<T> std::ops::Mul<T> for NdVector<#n, T> 
90            where 
91                T: std::ops::Mul<Output = T> + std::ops::MulAssign + Copy
92        {
93            type Output = NdVector<#n, T>;
94
95            fn mul(mut self, rhs: T) -> Self::Output {
96                self *= rhs;
97                self
98            }
99        }
100
101        impl<T> std::ops::MulAssign<T> for NdVector<#n, T> 
102        where 
103            T: std::ops::MulAssign + Copy
104        {
105            fn mul_assign(&mut self, rhs: T){
106                unsafe { #(#indices_mul_assign)* }
107            }
108        }
109
110        impl<T> std::ops::Div<T> for NdVector<#n, T> 
111            where 
112                T: std::ops::Div<Output = T> + std::ops::DivAssign + Copy
113        {
114            type Output = NdVector<#n, T>;
115
116            fn div(mut self, rhs: T) -> Self::Output {
117                self /= rhs;
118                self
119            }
120        }
121
122        impl<T> std::ops::DivAssign<T> for NdVector<#n, T> 
123        where 
124            T: std::ops::DivAssign + Copy
125        {
126            fn div_assign(&mut self, rhs: T){
127                unsafe { #(#indices_div_assign)* }
128            }
129        }
130
131
132        impl<Data> Dot for NdVector<#n, Data>
133            where Data: DataValue 
134        {
135            type Product = Data;
136            fn dot(self, rhs: Self) -> Self::Product {
137                let mut result = Data::zero();
138                unsafe {
139                    #(#indices_dot)*
140                };
141                result
142            }
143        }
144
145
146
147        impl<Data> ElementWiseMul<Self> for NdVector<#n, Data> 
148            where Data: DataValue + ops::MulAssign
149        {
150            type Output = Self;
151            fn elem_mul(mut self, rhs: Self) -> Self::Output {
152                unsafe { #(#indices_elem_mul)* }
153                self
154            }
155        }
156
157        impl<Data> ElementWiseDiv<Self> for NdVector<#n, Data> 
158            where Data: DataValue + ops::DivAssign
159        {
160            type Output = Self;
161            fn elem_div(mut self, rhs: Self) -> Self::Output {
162                unsafe { #(#indices_elem_div)* }
163                self
164            }
165        }
166
167        impl<Data> cmp::PartialEq for NdVector<#n, Data> 
168            where Data: PartialEq
169        {
170            fn eq(&self, rhs: &Self) -> bool {
171                let mut result = true;
172                unsafe { #(#indices_partial_eq)* }
173                result
174            }
175        }
176
177        impl<Data> Portable for NdVector<#n, Data> 
178            where Data: DataValue
179        {
180            fn to_bytes(self) -> Vec<u8> {
181                let mut result = Vec::with_capacity(#n * size_of::<Data>());
182                unsafe { #(#indices_portable_to_bytes)* }
183                result
184            }
185            
186            fn write_to<W>(self, writer: &mut W) 
187                where W: ByteWriter
188            {
189                unsafe{ #(#indices_portable_write_to)* }
190            }
191
192            fn read_from<R>(reader: &mut R) -> Result<Self, ReaderErr>
193                where R: ByteReader
194            {
195                let mut data = [Data::zero(); #n];
196                unsafe { #(#indices_portable_read_from)* }
197                Ok(Self {
198                    data,
199                })
200            }
201        }
202
203
204        macro_rules! impl_vector {
205            ($($t:ty);* ) => {
206            $(
207                impl Vector<#n> for NdVector<#n, $t> 
208                {
209                    type Component = $t;
210                    fn zero() -> Self {
211                        Self {
212                            data: [<$t as DataValue>::zero(); #n],
213                        }
214                    }
215                    fn get(&self, index: usize) -> &Self::Component {
216                        self.data.index(index)
217                    }
218                    
219                    fn get_mut(&mut self, index: usize) -> &mut Self::Component {
220                        self.data.index_mut(index)
221                    }
222                    
223                    unsafe fn get_unchecked(&self, index: usize) -> &Self::Component {
224                        self.data.as_slice().get_unchecked(index)
225                    }
226                    
227                    unsafe fn get_unchecked_mut(&mut self, index: usize) -> &mut Self::Component {
228                        self.data.as_mut_slice().get_unchecked_mut(index)
229                    }
230                }
231            )*
232            };
233        }
234
235        impl_vector!(
236            u8; u16; u32; u64; i8; i16; i32; i64;
237            f32; f64
238        );
239    };
240
241    TokenStream::from(expanded)
242}
243
244
245
246#[proc_macro]
247pub fn impl_ndvector_ops(_input: TokenStream) -> TokenStream {
248    let n: usize = env::var("MAX_VECTOR_DIM")
249        .ok()
250        .and_then(|s| s.parse().ok())
251        .unwrap_or(4); // Default max vec size is 4 if missing or invalid
252
253    let expanded = quote! {
254        use draco_nd_vector::impl_ndvector_ops_for_dim;
255        seq_macro::seq!(N in 1..=#n {
256            impl_ndvector_ops_for_dim!(N);
257        });
258    };
259
260    TokenStream::from(expanded)
261}