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); 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}