Skip to main content

nir_rs/
nodes.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2
3//! Wire-accurate NIR computational node types.
4//!
5//! The closed [`NirNode`] enum mirrors upstream HDF5 `type` strings exactly
6//! (`CubaLIF`, `Conv2d`, `SumPool2d`, `I`, …). Do **not** invent marketing
7//! aliases (`CurrLIF`, `Convolution`, `Integrator`) — those break
8//! interoperability with Python NIR.
9//!
10//! Field names use snake_case matching the neuromorphs/NIR Python dataclasses.
11//! Numeric array parameters are [`Tensor`] values (Python: `numpy.ndarray`).
12
13use crate::graph::NirGraph;
14use crate::types::{MetadataMap, Tensor};
15
16/// Convolution / pooling padding specification.
17///
18/// Upstream NIR accepts integer extents or the string modes `"same"` / `"valid"`.
19#[derive(Debug, Clone, PartialEq)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21#[non_exhaustive]
22pub enum Padding {
23    /// Explicit per-axis padding extents (length 1 for 1d, 2 for 2d, …).
24    Explicit(Vec<i64>),
25    /// Pad so that spatial output size matches input (`"same"`).
26    Same,
27    /// No padding (`"valid"`).
28    Valid,
29}
30
31impl Padding {
32    /// Single-axis explicit padding.
33    #[must_use]
34    pub fn single(value: i64) -> Self {
35        Self::Explicit(vec![value])
36    }
37
38    /// Two-axis explicit padding `(h, w)`.
39    #[must_use]
40    pub fn pair(h: i64, w: i64) -> Self {
41        Self::Explicit(vec![h, w])
42    }
43}
44
45/// Closed set of NIR computational nodes.
46///
47/// Exhaustive matching is intentional so downstream tools can cover every wire type.
48/// This enum is **not** `#[non_exhaustive]` so downstream mappers can cover all
49/// wire types without a wildcard arm (new wire types are a major API change).
50/// With the `serde` feature, the representation is internally tagged by
51/// `"type"`; every tag is explicitly renamed to [`Self::type_name`].
52#[derive(Debug, Clone, PartialEq)]
53#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
54#[cfg_attr(feature = "serde", serde(tag = "type"))]
55pub enum NirNode {
56    /// Graph input port (`type = "Input"`).
57    #[cfg_attr(feature = "serde", serde(rename = "Input"))]
58    Input(Input),
59    /// Graph output port (`type = "Output"`).
60    #[cfg_attr(feature = "serde", serde(rename = "Output"))]
61    Output(Output),
62    /// Affine transform `y = W x + b` (`type = "Affine"`).
63    #[cfg_attr(feature = "serde", serde(rename = "Affine"))]
64    Affine(Affine),
65    /// Linear transform without bias (`type = "Linear"`).
66    #[cfg_attr(feature = "serde", serde(rename = "Linear"))]
67    Linear(Linear),
68    /// Elementwise scale (`type = "Scale"`).
69    #[cfg_attr(feature = "serde", serde(rename = "Scale"))]
70    Scale(Scale),
71    /// 1-D convolution (`type = "Conv1d"`).
72    #[cfg_attr(feature = "serde", serde(rename = "Conv1d"))]
73    Conv1d(Conv1d),
74    /// 2-D convolution (`type = "Conv2d"`).
75    #[cfg_attr(feature = "serde", serde(rename = "Conv2d"))]
76    Conv2d(Conv2d),
77    /// Current-based leaky integrator (`type = "CubaLI"`).
78    #[cfg_attr(feature = "serde", serde(rename = "CubaLI"))]
79    CubaLi(CubaLi),
80    /// Current-based LIF (`type = "CubaLIF"`).
81    #[cfg_attr(feature = "serde", serde(rename = "CubaLIF"))]
82    CubaLif(CubaLif),
83    /// Pure delay (`type = "Delay"`).
84    #[cfg_attr(feature = "serde", serde(rename = "Delay"))]
85    Delay(Delay),
86    /// Flatten (`type = "Flatten"`).
87    #[cfg_attr(feature = "serde", serde(rename = "Flatten"))]
88    Flatten(Flatten),
89    /// Integrator (`type = "I"`).
90    #[cfg_attr(feature = "serde", serde(rename = "I"))]
91    I(I),
92    /// Integrate-and-fire (`type = "IF"`).
93    #[cfg_attr(feature = "serde", serde(rename = "IF"))]
94    If(If),
95    /// Leaky integrator (`type = "LI"`).
96    #[cfg_attr(feature = "serde", serde(rename = "LI"))]
97    Li(Li),
98    /// Leaky integrate-and-fire (`type = "LIF"`).
99    #[cfg_attr(feature = "serde", serde(rename = "LIF"))]
100    Lif(Lif),
101    /// Sum pooling 2-D (`type = "SumPool2d"`).
102    #[cfg_attr(feature = "serde", serde(rename = "SumPool2d"))]
103    SumPool2d(SumPool2d),
104    /// Average pooling 2-D (`type = "AvgPool2d"`).
105    #[cfg_attr(feature = "serde", serde(rename = "AvgPool2d"))]
106    AvgPool2d(AvgPool2d),
107    /// Heaviside threshold (`type = "Threshold"`).
108    #[cfg_attr(feature = "serde", serde(rename = "Threshold"))]
109    Threshold(Threshold),
110    /// Nested subgraph (`type = "NIRGraph"`).
111    #[cfg_attr(feature = "serde", serde(rename = "NIRGraph"))]
112    Graph(Box<NirGraph>),
113}
114
115impl NirNode {
116    /// Exact upstream HDF5 / Python wire `type` string.
117    #[must_use]
118    pub fn type_name(&self) -> &'static str {
119        match self {
120            Self::Input(_) => "Input",
121            Self::Output(_) => "Output",
122            Self::Affine(_) => "Affine",
123            Self::Linear(_) => "Linear",
124            Self::Scale(_) => "Scale",
125            Self::Conv1d(_) => "Conv1d",
126            Self::Conv2d(_) => "Conv2d",
127            Self::CubaLi(_) => "CubaLI",
128            Self::CubaLif(_) => "CubaLIF",
129            Self::Delay(_) => "Delay",
130            Self::Flatten(_) => "Flatten",
131            Self::I(_) => "I",
132            Self::If(_) => "IF",
133            Self::Li(_) => "LI",
134            Self::Lif(_) => "LIF",
135            Self::SumPool2d(_) => "SumPool2d",
136            Self::AvgPool2d(_) => "AvgPool2d",
137            Self::Threshold(_) => "Threshold",
138            Self::Graph(_) => "NIRGraph",
139        }
140    }
141
142    /// Validate local convolution and pooling parameter invariants.
143    ///
144    /// Nodes other than [`Self::Conv1d`], [`Self::Conv2d`], [`Self::SumPool2d`],
145    /// [`Self::AvgPool2d`], and nested [`Self::Graph`] succeed without further
146    /// checks. This does **not** run shape inference or neuron-parameter
147    /// validation.
148    ///
149    /// Isolated calls report failures against the node name `"<node>"`. Prefer
150    /// [`crate::NirGraph::validate_parameters`] when the node lives in a graph
151    /// so the error names the map key.
152    ///
153    /// # Errors
154    ///
155    /// [`crate::NirError::InvalidNodeParameters`] when a convolution or pooling
156    /// field violates a local invariant. Nested subgraphs are visited.
157    pub fn validate_parameters(&self) -> crate::error::Result<()> {
158        crate::validation::validate_node(self, crate::validation::ANONYMOUS_NODE)
159    }
160}
161
162/// Input port: virtual node feeding data into the graph.
163///
164/// Wire field: `shape` (array of axis lengths).
165#[derive(Debug, Clone, PartialEq)]
166#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
167pub struct Input {
168    /// Shape of the input tensor.
169    pub shape: Vec<usize>,
170    /// Free-form node metadata (Python `metadata` dict).
171    pub metadata: MetadataMap,
172}
173
174/// Output port: virtual node collecting graph results.
175///
176/// Wire field: `shape` (array of axis lengths).
177#[derive(Debug, Clone, PartialEq)]
178#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
179pub struct Output {
180    /// Shape of the output tensor.
181    pub shape: Vec<usize>,
182    /// Free-form node metadata (Python `metadata` dict).
183    pub metadata: MetadataMap,
184}
185
186/// Affine map `y = W x + b`.
187#[derive(Debug, Clone, PartialEq)]
188#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
189pub struct Affine {
190    /// Weight matrix / tensor.
191    pub weight: Tensor,
192    /// Bias vector / tensor.
193    pub bias: Tensor,
194    /// Free-form node metadata (Python `metadata` dict).
195    pub metadata: MetadataMap,
196}
197
198/// Linear map without bias `y = W x`.
199#[derive(Debug, Clone, PartialEq)]
200#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
201pub struct Linear {
202    /// Weight matrix / tensor.
203    pub weight: Tensor,
204    /// Free-form node metadata (Python `metadata` dict).
205    pub metadata: MetadataMap,
206}
207
208/// Elementwise scale `y = x ⊙ s`.
209#[derive(Debug, Clone, PartialEq)]
210#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
211pub struct Scale {
212    /// Per-element scale factors.
213    pub scale: Tensor,
214    /// Free-form node metadata (Python `metadata` dict).
215    pub metadata: MetadataMap,
216}
217
218/// 1-D convolution.
219#[derive(Debug, Clone, PartialEq)]
220#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
221pub struct Conv1d {
222    /// Kernel weights, typically `(C_out, C_in, K)`.
223    pub weight: Tensor,
224    /// Stride (scalar or length-1).
225    pub stride: Vec<i64>,
226    /// Padding specification.
227    pub padding: Padding,
228    /// Dilation (scalar or length-1).
229    pub dilation: Vec<i64>,
230    /// Grouped convolution groups.
231    pub groups: i64,
232    /// Bias of shape `(C_out,)` (required on the NIR wire).
233    pub bias: Tensor,
234    /// Optional spatial input length `N` used for shape inference.
235    pub input_shape: Option<usize>,
236    /// Free-form node metadata (Python `metadata` dict).
237    pub metadata: MetadataMap,
238}
239
240/// 2-D convolution.
241#[derive(Debug, Clone, PartialEq)]
242#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
243pub struct Conv2d {
244    /// Kernel weights, typically `(C_out, C_in, Kh, Kw)`.
245    pub weight: Tensor,
246    /// Stride per spatial axis (or single value expanded later).
247    pub stride: Vec<i64>,
248    /// Padding specification.
249    pub padding: Padding,
250    /// Dilation per spatial axis.
251    pub dilation: Vec<i64>,
252    /// Grouped convolution groups.
253    pub groups: i64,
254    /// Bias of shape `(C_out,)` (required on the NIR wire).
255    pub bias: Tensor,
256    /// Optional spatial input `(N_x, N_y)`.
257    pub input_shape: Option<Vec<usize>>,
258    /// Free-form node metadata (Python `metadata` dict).
259    pub metadata: MetadataMap,
260}
261
262/// Current-based leaky integrator (`CubaLI`).
263#[derive(Debug, Clone, PartialEq)]
264#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
265pub struct CubaLi {
266    /// Synaptic time constant.
267    pub tau_syn: Tensor,
268    /// Membrane time constant.
269    pub tau_mem: Tensor,
270    /// Resistance.
271    pub r: Tensor,
272    /// Leak voltage.
273    pub v_leak: Tensor,
274    /// Input current weight (elementwise).
275    ///
276    /// Upstream Python NIR defaults missing `w_in` to ones (broadcast). Use
277    /// [`None`] when the field is absent on the wire; v0.3 decode should
278    /// synthesize ones when needed.
279    pub w_in: Option<Tensor>,
280    /// Free-form node metadata (Python `metadata` dict).
281    pub metadata: MetadataMap,
282}
283
284/// Current-based leaky integrate-and-fire (`CubaLIF`).
285#[derive(Debug, Clone, PartialEq)]
286#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
287pub struct CubaLif {
288    /// Synaptic time constant.
289    pub tau_syn: Tensor,
290    /// Membrane time constant.
291    pub tau_mem: Tensor,
292    /// Resistance.
293    pub r: Tensor,
294    /// Leak voltage.
295    pub v_leak: Tensor,
296    /// Firing threshold.
297    pub v_threshold: Tensor,
298    /// Reset potential (optional; Python defaults to zeros).
299    pub v_reset: Option<Tensor>,
300    /// Input current weight (elementwise).
301    ///
302    /// Upstream Python NIR defaults missing `w_in` to ones (broadcast). Use
303    /// [`None`] when the field is absent on the wire; v0.3 decode should
304    /// synthesize ones when needed.
305    pub w_in: Option<Tensor>,
306    /// Free-form node metadata (Python `metadata` dict).
307    pub metadata: MetadataMap,
308}
309
310/// Pure delay `y(t) = x(t − τ)`.
311#[derive(Debug, Clone, PartialEq)]
312#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
313pub struct Delay {
314    /// Delay amount(s).
315    pub delay: Tensor,
316    /// Free-form node metadata (Python `metadata` dict).
317    pub metadata: MetadataMap,
318}
319
320/// Flatten a contiguous range of dimensions.
321#[derive(Debug, Clone, PartialEq)]
322#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
323pub struct Flatten {
324    /// First dimension to flatten (Python default: 1).
325    pub start_dim: i64,
326    /// Last dimension to flatten (Python default: −1).
327    pub end_dim: i64,
328    /// Optional input shape used for shape inference / wire `input_type`.
329    pub input_type: Option<Vec<usize>>,
330    /// Free-form node metadata (Python `metadata` dict).
331    pub metadata: MetadataMap,
332}
333
334/// Integrator neuron (`I`): `dv/dt = R I`.
335#[derive(Debug, Clone, PartialEq)]
336#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
337pub struct I {
338    /// Resistance.
339    pub r: Tensor,
340    /// Free-form node metadata (Python `metadata` dict).
341    pub metadata: MetadataMap,
342}
343
344/// Integrate-and-fire neuron (`IF`).
345#[derive(Debug, Clone, PartialEq)]
346#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
347pub struct If {
348    /// Resistance.
349    pub r: Tensor,
350    /// Firing threshold.
351    pub v_threshold: Tensor,
352    /// Reset potential (optional).
353    pub v_reset: Option<Tensor>,
354    /// Free-form node metadata (Python `metadata` dict).
355    pub metadata: MetadataMap,
356}
357
358/// Leaky integrator (`LI`).
359#[derive(Debug, Clone, PartialEq)]
360#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
361pub struct Li {
362    /// Membrane time constant.
363    pub tau: Tensor,
364    /// Resistance.
365    pub r: Tensor,
366    /// Leak voltage.
367    pub v_leak: Tensor,
368    /// Free-form node metadata (Python `metadata` dict).
369    pub metadata: MetadataMap,
370}
371
372/// Leaky integrate-and-fire (`LIF`).
373#[derive(Debug, Clone, PartialEq)]
374#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
375pub struct Lif {
376    /// Membrane time constant.
377    pub tau: Tensor,
378    /// Resistance.
379    pub r: Tensor,
380    /// Leak voltage.
381    pub v_leak: Tensor,
382    /// Firing threshold.
383    pub v_threshold: Tensor,
384    /// Reset potential (optional).
385    pub v_reset: Option<Tensor>,
386    /// Free-form node metadata (Python `metadata` dict).
387    pub metadata: MetadataMap,
388}
389
390/// 2-D sum pooling.
391#[derive(Debug, Clone, PartialEq)]
392#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
393pub struct SumPool2d {
394    /// Kernel size `(H, W)`.
395    pub kernel_size: Tensor,
396    /// Stride `(H, W)`.
397    pub stride: Tensor,
398    /// Padding `(H, W)`.
399    pub padding: Tensor,
400    /// Free-form node metadata (Python `metadata` dict).
401    pub metadata: MetadataMap,
402}
403
404/// 2-D average pooling.
405#[derive(Debug, Clone, PartialEq)]
406#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
407pub struct AvgPool2d {
408    /// Kernel size `(H, W)`.
409    pub kernel_size: Tensor,
410    /// Stride `(H, W)`.
411    pub stride: Tensor,
412    /// Padding `(H, W)`.
413    pub padding: Tensor,
414    /// Free-form node metadata (Python `metadata` dict).
415    pub metadata: MetadataMap,
416}
417
418/// Heaviside threshold / surrogate step.
419#[derive(Debug, Clone, PartialEq)]
420#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
421pub struct Threshold {
422    /// Threshold value(s).
423    pub threshold: Tensor,
424    /// Free-form node metadata (Python `metadata` dict).
425    pub metadata: MetadataMap,
426}
427
428#[cfg(test)]
429mod tests {
430    use super::*;
431    use crate::types::Tensor;
432
433    fn sample_weight() -> Tensor {
434        Tensor::from_f32(vec![2, 3], vec![1., 0., 0., 0., 1., 0.]).unwrap()
435    }
436
437    fn sample_bias() -> Tensor {
438        Tensor::from_f32(vec![2], vec![0., 0.]).unwrap()
439    }
440
441    fn sample_vec3() -> Tensor {
442        Tensor::from_f64(vec![3], vec![1.0, 1.0, 1.0]).unwrap()
443    }
444
445    /// Shared `(kernel_size, stride, padding)` tensors for SumPool2d / AvgPool2d
446    /// wire-name samples (keeps the two cases from being pure copy-paste).
447    fn sample_pool2d_fields() -> (Tensor, Tensor, Tensor) {
448        (
449            Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
450            Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
451            Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
452        )
453    }
454
455    #[test]
456    fn all_type_names_match_wire_strings() {
457        let cases: Vec<(&str, NirNode)> = vec![
458            (
459                "Input",
460                NirNode::Input(Input {
461                    shape: vec![1, 4],
462                    metadata: Default::default(),
463                }),
464            ),
465            (
466                "Output",
467                NirNode::Output(Output {
468                    shape: vec![1, 2],
469                    metadata: Default::default(),
470                }),
471            ),
472            (
473                "Affine",
474                NirNode::Affine(Affine {
475                    weight: sample_weight(),
476                    bias: sample_bias(),
477                    metadata: Default::default(),
478                }),
479            ),
480            (
481                "Linear",
482                NirNode::Linear(Linear {
483                    weight: sample_weight(),
484                    metadata: Default::default(),
485                }),
486            ),
487            (
488                "Scale",
489                NirNode::Scale(Scale {
490                    scale: sample_vec3(),
491                    metadata: Default::default(),
492                }),
493            ),
494            (
495                "Conv1d",
496                NirNode::Conv1d(Conv1d {
497                    weight: Tensor::from_f32(vec![1, 1, 3], vec![1., 0., -1.]).unwrap(),
498                    stride: vec![1],
499                    padding: Padding::single(0),
500                    dilation: vec![1],
501                    groups: 1,
502                    bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
503                    input_shape: Some(10),
504                    metadata: Default::default(),
505                }),
506            ),
507            (
508                "Conv2d",
509                NirNode::Conv2d(Conv2d {
510                    weight: Tensor::from_f32(vec![1, 1, 3, 3], vec![0.; 9]).unwrap(),
511                    stride: vec![1, 1],
512                    padding: Padding::Same,
513                    dilation: vec![1, 1],
514                    groups: 1,
515                    bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
516                    input_shape: Some(vec![28, 28]),
517                    metadata: Default::default(),
518                }),
519            ),
520            (
521                "CubaLI",
522                NirNode::CubaLi(CubaLi {
523                    tau_syn: sample_vec3(),
524                    tau_mem: sample_vec3(),
525                    r: sample_vec3(),
526                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
527                    w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
528                    metadata: Default::default(),
529                }),
530            ),
531            (
532                "CubaLIF",
533                NirNode::CubaLif(CubaLif {
534                    tau_syn: sample_vec3(),
535                    tau_mem: sample_vec3(),
536                    r: sample_vec3(),
537                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
538                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
539                    v_reset: None,
540                    w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
541                    metadata: Default::default(),
542                }),
543            ),
544            (
545                "Delay",
546                NirNode::Delay(Delay {
547                    delay: Tensor::scalar_f64(1.0),
548                    metadata: Default::default(),
549                }),
550            ),
551            (
552                "Flatten",
553                NirNode::Flatten(Flatten {
554                    start_dim: 1,
555                    end_dim: -1,
556                    input_type: Some(vec![1, 4, 4]),
557                    metadata: Default::default(),
558                }),
559            ),
560            (
561                "I",
562                NirNode::I(I {
563                    r: sample_vec3(),
564                    metadata: Default::default(),
565                }),
566            ),
567            (
568                "IF",
569                NirNode::If(If {
570                    r: sample_vec3(),
571                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
572                    v_reset: None,
573                    metadata: Default::default(),
574                }),
575            ),
576            (
577                "LI",
578                NirNode::Li(Li {
579                    tau: sample_vec3(),
580                    r: sample_vec3(),
581                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
582                    metadata: Default::default(),
583                }),
584            ),
585            (
586                "LIF",
587                NirNode::Lif(Lif {
588                    tau: sample_vec3(),
589                    r: sample_vec3(),
590                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
591                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
592                    v_reset: Some(Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap()),
593                    metadata: Default::default(),
594                }),
595            ),
596            {
597                let (kernel_size, stride, padding) = sample_pool2d_fields();
598                (
599                    "SumPool2d",
600                    NirNode::SumPool2d(SumPool2d {
601                        kernel_size,
602                        stride,
603                        padding,
604                        metadata: Default::default(),
605                    }),
606                )
607            },
608            {
609                let (kernel_size, stride, padding) = sample_pool2d_fields();
610                (
611                    "AvgPool2d",
612                    NirNode::AvgPool2d(AvgPool2d {
613                        kernel_size,
614                        stride,
615                        padding,
616                        metadata: Default::default(),
617                    }),
618                )
619            },
620            (
621                "Threshold",
622                NirNode::Threshold(Threshold {
623                    threshold: Tensor::scalar_f64(1.0),
624                    metadata: Default::default(),
625                }),
626            ),
627            ("NIRGraph", NirNode::Graph(Box::new(NirGraph::new()))),
628        ];
629
630        assert_eq!(cases.len(), 19, "expected all wire node types");
631        for (wire, node) in cases {
632            assert_eq!(node.type_name(), wire);
633            #[cfg(feature = "serde")]
634            {
635                let value = serde_json::to_value(&node).unwrap();
636                assert_eq!(value["type"], wire);
637                assert_eq!(serde_json::from_value::<NirNode>(value).unwrap(), node);
638            }
639        }
640    }
641
642    #[test]
643    fn padding_helpers() {
644        assert_eq!(Padding::single(1), Padding::Explicit(vec![1]));
645        assert_eq!(Padding::pair(1, 2), Padding::Explicit(vec![1, 2]));
646        let _ = Padding::Same;
647        let _ = Padding::Valid;
648    }
649
650    #[test]
651    fn never_use_marketing_aliases() {
652        // Guard against accidental CurrLIF / Convolution / Integrator names.
653        let names: Vec<&str> = [
654            NirNode::CubaLif(CubaLif {
655                tau_syn: Tensor::scalar_f64(1.0),
656                tau_mem: Tensor::scalar_f64(1.0),
657                r: Tensor::scalar_f64(1.0),
658                v_leak: Tensor::scalar_f64(0.0),
659                v_threshold: Tensor::scalar_f64(1.0),
660                v_reset: None,
661                w_in: Some(Tensor::scalar_f64(1.0)),
662                metadata: Default::default(),
663            }),
664            NirNode::Conv2d(Conv2d {
665                weight: Tensor::from_f32(vec![1, 1, 1, 1], vec![1.]).unwrap(),
666                stride: vec![1, 1],
667                padding: Padding::Valid,
668                dilation: vec![1, 1],
669                groups: 1,
670                bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
671                input_shape: None,
672                metadata: Default::default(),
673            }),
674            NirNode::I(I {
675                r: Tensor::scalar_f64(1.0),
676                metadata: Default::default(),
677            }),
678            NirNode::SumPool2d(SumPool2d {
679                kernel_size: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
680                stride: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
681                padding: Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
682                metadata: Default::default(),
683            }),
684        ]
685        .into_iter()
686        .map(|n| n.type_name())
687        .collect();
688
689        assert_eq!(names, ["CubaLIF", "Conv2d", "I", "SumPool2d"]);
690        for n in names {
691            assert!(!n.contains("Curr"));
692            assert!(!n.contains("Convolution"));
693            assert!(!n.contains("Integrator"));
694            assert!(!n.contains("Pooling"));
695        }
696    }
697}