1#![forbid(unsafe_code)]
31#![doc = include_str!("../README.md")]
32#![forbid(rustdoc::broken_intra_doc_links)]
33#![forbid(rustdoc::private_intra_doc_links)]
34#![forbid(missing_docs)]
35#![forbid(rustdoc::missing_crate_level_docs)]
36#![forbid(rustdoc::private_doc_tests)]
38#![forbid(rustdoc::invalid_codeblock_attributes)]
39#![forbid(rustdoc::invalid_html_tags)]
40#![forbid(rustdoc::invalid_rust_codeblocks)]
41#![forbid(rustdoc::bare_urls)]
42#![forbid(rustdoc::unescaped_backticks)]
43#![forbid(rustdoc::redundant_explicit_links)]
44
45use proc_macro::TokenStream;
46use quote::quote;
47use syn::{parse_macro_input, Data, DataStruct, DeriveInput};
48
49#[proc_macro_derive(Module, attributes(no_param))]
58pub fn module_derive(input: TokenStream) -> TokenStream {
59 let input = parse_macro_input!(input as DeriveInput);
60 derive_module(&input)
61}
62
63fn derive_module(input: &DeriveInput) -> TokenStream {
64 let struct_name = &input.ident;
65 let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
66
67 let mut field_iterators = quote! {
69 trait __MarkerTraitRef: Sized {
70 fn __iterate_by_ref(self, res: &mut Vec<(String, &zyx::Tensor)>, label: &str) {}
71 }
72
73 struct __MarkerStructRef<T>(T);
74
75 impl<'a, T: zyx::Module> __MarkerStructRef<&'a T> {
76 fn __iterate_by_ref(self, res: &mut Vec<(String, &'a zyx::Tensor)>, label: &str) {
77 res.extend(self.0.iter_tensors().map(|(k, t)| (format!("{label}.{k}"), t)));
78 }
79 }
80
81 impl<'a, T> __MarkerTraitRef for __MarkerStructRef<&'a T>{}
82
83 let mut res = Vec::<(String, &zyx::Tensor)>::new();
84 };
85
86 if let Data::Struct(DataStruct { fields, .. }) = &input.data {
87 for field in fields.iter() {
88 let field_name = match &field.ident {
89 Some(ident) => ident,
90 None => panic!("Unnamed fields are not supported"),
91 };
92 let field_name_str = field_name.to_string();
93
94 let field_ty: &syn::Type = &field.ty;
95 let no_param = has_no_param_attr(&field.attrs);
96
97 use std::string::ToString;
98 if no_param {
99 } else if quote! { #field_ty }.to_string() == "Tensor" {
101 field_iterators = quote! {
102 #field_iterators
103 res.push((#field_name_str.to_string(), &self.#field_name));
104 }
105 } else if quote! { #field_ty }.to_string() == "Option < Tensor >" {
106 field_iterators = quote! {
107 #field_iterators
108 if let Some(tensor) = &self.#field_name {
109 res.push((#field_name_str.to_string(), tensor));
110 }
111 }
112 } else {
113 field_iterators = quote! {
114 #field_iterators
115 __MarkerStructRef::<&#field_ty>::__iterate_by_ref(__MarkerStructRef(&self.#field_name), &mut res, #field_name_str);
116 };
117 }
118 }
119 }
120
121 let mut mut_field_iterators = quote! {
123 trait __MarkerTraitRef: Sized {
124 fn __iterate_by_ref(mut self, res: &mut Vec<(String, &mut zyx::Tensor)>, label: &str) {}
125 }
126
127 struct __MarkerStructRef<T>(T);
128
129 impl<'a, T: zyx::Module> __MarkerStructRef<&'a mut T> {
130 fn __iterate_by_ref(mut self, res: &mut Vec<(String, &'a mut zyx::Tensor)>, label: &str) {
131 res.extend(self.0.iter_tensors_mut().map(|(k, t)| (format!("{label}.{k}"), t)));
132 }
133 }
134
135 impl<'a, T> __MarkerTraitRef for __MarkerStructRef<&'a mut T>{}
136
137 let mut res = Vec::<(String, &mut zyx::Tensor)>::new();
138 };
139
140 if let Data::Struct(DataStruct { fields, .. }) = &input.data {
141 for field in fields.iter() {
142 let field_name = match &field.ident {
143 Some(ident) => ident,
144 None => panic!("Unnamed fields are not supported"),
145 };
146 let field_name_str = field_name.to_string();
147
148 let field_ty: &syn::Type = &field.ty;
149 let no_param = has_no_param_attr(&field.attrs);
150
151 use std::string::ToString;
152 if no_param {
153 } else if quote! { #field_ty }.to_string() == "Tensor" {
155 mut_field_iterators = quote! {
156 #mut_field_iterators
157 res.push((#field_name_str.to_string(), &mut self.#field_name));
158 }
159 } else if quote! { #field_ty }.to_string() == "Option < Tensor >" {
160 mut_field_iterators = quote! {
161 #mut_field_iterators
162 if let Some(tensor) = &mut self.#field_name {
163 res.push((#field_name_str.to_string(), tensor));
164 }
165 }
166 } else {
167 mut_field_iterators = quote! {
168 #mut_field_iterators
169 __MarkerStructRef::<&mut #field_ty>::__iterate_by_ref(__MarkerStructRef(&mut self.#field_name), &mut res, #field_name_str);
170 };
171 }
172 }
173 }
174
175 let expanded = quote! {
176 impl #impl_generics zyx::Module for #struct_name #ty_generics #where_clause {
177 fn iter<'a>(&'a self) -> impl Iterator<Item = &'a zyx::Tensor> {
178 self.into_iter()
179 }
180
181 fn iter_mut<'a>(&'a mut self) -> impl Iterator<Item = &'a mut zyx::Tensor> {
182 self.into_iter()
183 }
184
185 fn iter_tensors<'a>(&'a self) -> impl Iterator<Item = (String, &'a zyx::Tensor)> {
186 #field_iterators
187 res.into_iter()
188 }
189
190 fn iter_tensors_mut<'a>(&'a mut self) -> impl Iterator<Item = (String, &'a mut zyx::Tensor)> {
191 #mut_field_iterators
192 res.into_iter()
193 }
194 }
195 };
196
197 let mut field_iterators = quote! {
199 trait __MarkerTraitRef<'a> {
200 fn __iterate_by_ref(&self, res: &mut Vec<&'a zyx::Tensor>) {}
201 }
202
203 struct __MarkerStructRef<T: Copy>(T);
204
205 impl<'a, T: IntoIterator<Item = &'a zyx::Tensor> + Copy> __MarkerStructRef<T> {
206 fn __iterate_by_ref(&self, res: &mut Vec<&'a zyx::Tensor>) {
207 res.extend(self.0.into_iter());
208 }
209 }
210
211 impl<'a, T: Copy> __MarkerTraitRef<'a> for __MarkerStructRef<T>{}
212
213 let mut res = Vec::<&zyx::Tensor>::new();
214 };
215
216 if let Data::Struct(DataStruct { fields, .. }) = &input.data {
217 for field in fields.iter() {
218 let field_name = match &field.ident {
219 Some(ident) => ident,
220 None => panic!("Unnamed fields are not supported"),
221 };
222 let field_ty: &syn::Type = &field.ty;
223 let no_param = has_no_param_attr(&field.attrs);
224 use std::string::ToString;
225 if no_param {
226 } else if quote! { #field_ty }.to_string() == "Tensor" {
228 field_iterators = quote! {
229 #field_iterators
230 res.push(&self.#field_name);
231 }
232 } else {
233 field_iterators = quote! {
234 #field_iterators
235 __MarkerStructRef::<&#field_ty>::__iterate_by_ref(&__MarkerStructRef(&self.#field_name), &mut res);
236 };
237 }
238 }
239 }
240
241 let expanded = quote! {
242 #expanded
243
244 impl<'a> IntoIterator for &'a #struct_name #ty_generics {
245 type Item = &'a zyx::Tensor;
246 type IntoIter = std::vec::IntoIter<&'a zyx::Tensor>;
247
248 fn into_iter(self) -> Self::IntoIter {
249 #field_iterators
250 res.into_iter()
251 }
252 }
253 };
254
255 let mut field_iterators = quote! {
257 trait MarkerTraitMut<'a>: Sized {
258 fn iterate_by_mut(mut self, res: &mut Vec<&'a mut zyx::Tensor>) {}
259 }
260
261 struct MarkerStructMut<T>(T);
262
263 impl<'a, T: IntoIterator<Item = &'a mut zyx::Tensor>> MarkerStructMut<T> {
264 fn iterate_by_mut(mut self, res: &mut Vec<&'a mut zyx::Tensor>) {
265 res.extend(self.0.into_iter());
266 }
267 }
268
269 impl<'a, T> MarkerTraitMut<'a> for MarkerStructMut<T>{}
270
271 let mut res = Vec::<&mut zyx::Tensor>::new();
272 };
273
274 if let Data::Struct(DataStruct { fields, .. }) = &input.data {
275 for field in fields.iter() {
276 let field_name = match &field.ident {
277 Some(ident) => ident,
278 None => panic!("Unnamed fields are not supported"),
279 };
280 let field_ty: &syn::Type = &field.ty;
281 let no_param = has_no_param_attr(&field.attrs);
282 use std::string::ToString;
283 if no_param {
284 } else if quote! { #field_ty }.to_string() == "Tensor" {
286 field_iterators = quote! {
287 #field_iterators
288 res.push(&mut self.#field_name);
289 }
290 } else {
291 field_iterators = quote! {
292 #field_iterators
293 MarkerStructMut::<&mut #field_ty>::iterate_by_mut(MarkerStructMut(&mut self.#field_name), &mut res);
294 };
295 }
296 }
297 }
298
299 let expanded = quote! {
300 #expanded
301
302 impl<'a> IntoIterator for &'a mut #struct_name #ty_generics {
303 type Item = &'a mut zyx::Tensor;
304 type IntoIter = std::vec::IntoIter<&'a mut zyx::Tensor>;
305
306 fn into_iter(self) -> Self::IntoIter {
307 #field_iterators
308 res.into_iter()
309 }
310 }
311 };
312
313 TokenStream::from(expanded)
314}
315
316fn has_no_param_attr(attrs: &[syn::Attribute]) -> bool {
320 attrs.iter().any(|a| a.path().is_ident("no_param"))
321}