Skip to main content

zyx_derive/
lib.rs

1// Copyright (C) 2025 zk4x
2// SPDX-License-Identifier: LGPL-3.0-only WITH Classpath-exception-2.0
3
4//! # zyx-derive
5//!
6//! This crate contains procedural macros for zyx.
7//!
8//! Macro Module automatically implements IntoIterator<Item = &Tensor>
9//! for your module, so that you can use it in backpropagation and save it to disk.
10//! ```rust
11//! use zyx::Tensor;
12//! use zyx_derive::Module;
13//!
14//! #[derive(Module)]
15//! struct MyNet {
16//!     b: Tensor,
17//!     w: Tensor,
18//! }
19//!
20//! impl MyNet {
21//!     fn forward(&self, x: &Tensor) -> Tensor {
22//!         x.dot(&self.w).unwrap() + &self.b
23//!     }
24//! }
25//! ```
26//!
27//! For README, quick tutorial and source code, please visit `<https://www.github.com/zk4x/zyx>`.
28//!
29//! For more details, there is a [book](https://www.github.com/zk4x/zyx/tree/main/zyx-book).
30#![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::missing_doc_code_examples)]
37#![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/// Implements FromIterator<Item = (String, Tensor)> and Module for your struct.
50///
51/// This allows saving, loading, backpropagation and updating your modules.
52///
53/// Recognised field attributes:
54/// - `#[no_param]` — marks a `Tensor` / `Option<Tensor>` field as a
55///   non-trainable hyperparameter (or state buffer). The parameter iteration
56///   skips it.
57#[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    // iter_tensors (immutable)
68    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                // Skip non-trainable hyperparameter entirely.
100            } 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    // iter_tensors_mut
122    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                // Skip non-trainable hyperparameter entirely.
154            } 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    // Implementation of IntoIterator<Item = &Tensor>
198    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                // Skip non-trainable hyperparameter entirely.
227            } 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    // Implementation of IntoIterator<Item = &mut Tensor>
256    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                // Skip non-trainable hyperparameter entirely.
285            } 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
316/// Returns true if any of the field's attributes is `#[no_param]`, which marks a
317/// `Tensor` / `Option<Tensor>` field as a non-trainable hyperparameter (or
318/// state buffer) that the `Module` derive's parameter iteration should skip.
319fn has_no_param_attr(attrs: &[syn::Attribute]) -> bool {
320    attrs.iter().any(|a| a.path().is_ident("no_param"))
321}