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
143/// Input port: virtual node feeding data into the graph.
144///
145/// Wire field: `shape` (array of axis lengths).
146#[derive(Debug, Clone, PartialEq)]
147#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
148pub struct Input {
149    /// Shape of the input tensor.
150    pub shape: Vec<usize>,
151    /// Free-form node metadata (Python `metadata` dict).
152    pub metadata: MetadataMap,
153}
154
155/// Output port: virtual node collecting graph results.
156///
157/// Wire field: `shape` (array of axis lengths).
158#[derive(Debug, Clone, PartialEq)]
159#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
160pub struct Output {
161    /// Shape of the output tensor.
162    pub shape: Vec<usize>,
163    /// Free-form node metadata (Python `metadata` dict).
164    pub metadata: MetadataMap,
165}
166
167/// Affine map `y = W x + b`.
168#[derive(Debug, Clone, PartialEq)]
169#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
170pub struct Affine {
171    /// Weight matrix / tensor.
172    pub weight: Tensor,
173    /// Bias vector / tensor.
174    pub bias: Tensor,
175    /// Free-form node metadata (Python `metadata` dict).
176    pub metadata: MetadataMap,
177}
178
179/// Linear map without bias `y = W x`.
180#[derive(Debug, Clone, PartialEq)]
181#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
182pub struct Linear {
183    /// Weight matrix / tensor.
184    pub weight: Tensor,
185    /// Free-form node metadata (Python `metadata` dict).
186    pub metadata: MetadataMap,
187}
188
189/// Elementwise scale `y = x ⊙ s`.
190#[derive(Debug, Clone, PartialEq)]
191#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
192pub struct Scale {
193    /// Per-element scale factors.
194    pub scale: Tensor,
195    /// Free-form node metadata (Python `metadata` dict).
196    pub metadata: MetadataMap,
197}
198
199/// 1-D convolution.
200#[derive(Debug, Clone, PartialEq)]
201#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
202pub struct Conv1d {
203    /// Kernel weights, typically `(C_out, C_in, K)`.
204    pub weight: Tensor,
205    /// Stride (scalar or length-1).
206    pub stride: Vec<i64>,
207    /// Padding specification.
208    pub padding: Padding,
209    /// Dilation (scalar or length-1).
210    pub dilation: Vec<i64>,
211    /// Grouped convolution groups.
212    pub groups: i64,
213    /// Bias of shape `(C_out,)` (required on the NIR wire).
214    pub bias: Tensor,
215    /// Optional spatial input length `N` used for shape inference.
216    pub input_shape: Option<usize>,
217    /// Free-form node metadata (Python `metadata` dict).
218    pub metadata: MetadataMap,
219}
220
221/// 2-D convolution.
222#[derive(Debug, Clone, PartialEq)]
223#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
224pub struct Conv2d {
225    /// Kernel weights, typically `(C_out, C_in, Kh, Kw)`.
226    pub weight: Tensor,
227    /// Stride per spatial axis (or single value expanded later).
228    pub stride: Vec<i64>,
229    /// Padding specification.
230    pub padding: Padding,
231    /// Dilation per spatial axis.
232    pub dilation: Vec<i64>,
233    /// Grouped convolution groups.
234    pub groups: i64,
235    /// Bias of shape `(C_out,)` (required on the NIR wire).
236    pub bias: Tensor,
237    /// Optional spatial input `(N_x, N_y)`.
238    pub input_shape: Option<Vec<usize>>,
239    /// Free-form node metadata (Python `metadata` dict).
240    pub metadata: MetadataMap,
241}
242
243/// Current-based leaky integrator (`CubaLI`).
244#[derive(Debug, Clone, PartialEq)]
245#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
246pub struct CubaLi {
247    /// Synaptic time constant.
248    pub tau_syn: Tensor,
249    /// Membrane time constant.
250    pub tau_mem: Tensor,
251    /// Resistance.
252    pub r: Tensor,
253    /// Leak voltage.
254    pub v_leak: Tensor,
255    /// Input current weight (elementwise).
256    ///
257    /// Upstream Python NIR defaults missing `w_in` to ones (broadcast). Use
258    /// [`None`] when the field is absent on the wire; v0.3 decode should
259    /// synthesize ones when needed.
260    pub w_in: Option<Tensor>,
261    /// Free-form node metadata (Python `metadata` dict).
262    pub metadata: MetadataMap,
263}
264
265/// Current-based leaky integrate-and-fire (`CubaLIF`).
266#[derive(Debug, Clone, PartialEq)]
267#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
268pub struct CubaLif {
269    /// Synaptic time constant.
270    pub tau_syn: Tensor,
271    /// Membrane time constant.
272    pub tau_mem: Tensor,
273    /// Resistance.
274    pub r: Tensor,
275    /// Leak voltage.
276    pub v_leak: Tensor,
277    /// Firing threshold.
278    pub v_threshold: Tensor,
279    /// Reset potential (optional; Python defaults to zeros).
280    pub v_reset: Option<Tensor>,
281    /// Input current weight (elementwise).
282    ///
283    /// Upstream Python NIR defaults missing `w_in` to ones (broadcast). Use
284    /// [`None`] when the field is absent on the wire; v0.3 decode should
285    /// synthesize ones when needed.
286    pub w_in: Option<Tensor>,
287    /// Free-form node metadata (Python `metadata` dict).
288    pub metadata: MetadataMap,
289}
290
291/// Pure delay `y(t) = x(t − τ)`.
292#[derive(Debug, Clone, PartialEq)]
293#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
294pub struct Delay {
295    /// Delay amount(s).
296    pub delay: Tensor,
297    /// Free-form node metadata (Python `metadata` dict).
298    pub metadata: MetadataMap,
299}
300
301/// Flatten a contiguous range of dimensions.
302#[derive(Debug, Clone, PartialEq)]
303#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
304pub struct Flatten {
305    /// First dimension to flatten (Python default: 1).
306    pub start_dim: i64,
307    /// Last dimension to flatten (Python default: −1).
308    pub end_dim: i64,
309    /// Optional input shape used for shape inference / wire `input_type`.
310    pub input_type: Option<Vec<usize>>,
311    /// Free-form node metadata (Python `metadata` dict).
312    pub metadata: MetadataMap,
313}
314
315/// Integrator neuron (`I`): `dv/dt = R I`.
316#[derive(Debug, Clone, PartialEq)]
317#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
318pub struct I {
319    /// Resistance.
320    pub r: Tensor,
321    /// Free-form node metadata (Python `metadata` dict).
322    pub metadata: MetadataMap,
323}
324
325/// Integrate-and-fire neuron (`IF`).
326#[derive(Debug, Clone, PartialEq)]
327#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
328pub struct If {
329    /// Resistance.
330    pub r: Tensor,
331    /// Firing threshold.
332    pub v_threshold: Tensor,
333    /// Reset potential (optional).
334    pub v_reset: Option<Tensor>,
335    /// Free-form node metadata (Python `metadata` dict).
336    pub metadata: MetadataMap,
337}
338
339/// Leaky integrator (`LI`).
340#[derive(Debug, Clone, PartialEq)]
341#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
342pub struct Li {
343    /// Membrane time constant.
344    pub tau: Tensor,
345    /// Resistance.
346    pub r: Tensor,
347    /// Leak voltage.
348    pub v_leak: Tensor,
349    /// Free-form node metadata (Python `metadata` dict).
350    pub metadata: MetadataMap,
351}
352
353/// Leaky integrate-and-fire (`LIF`).
354#[derive(Debug, Clone, PartialEq)]
355#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
356pub struct Lif {
357    /// Membrane time constant.
358    pub tau: Tensor,
359    /// Resistance.
360    pub r: Tensor,
361    /// Leak voltage.
362    pub v_leak: Tensor,
363    /// Firing threshold.
364    pub v_threshold: Tensor,
365    /// Reset potential (optional).
366    pub v_reset: Option<Tensor>,
367    /// Free-form node metadata (Python `metadata` dict).
368    pub metadata: MetadataMap,
369}
370
371/// 2-D sum pooling.
372#[derive(Debug, Clone, PartialEq)]
373#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
374pub struct SumPool2d {
375    /// Kernel size `(H, W)`.
376    pub kernel_size: Tensor,
377    /// Stride `(H, W)`.
378    pub stride: Tensor,
379    /// Padding `(H, W)`.
380    pub padding: Tensor,
381    /// Free-form node metadata (Python `metadata` dict).
382    pub metadata: MetadataMap,
383}
384
385/// 2-D average pooling.
386#[derive(Debug, Clone, PartialEq)]
387#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
388pub struct AvgPool2d {
389    /// Kernel size `(H, W)`.
390    pub kernel_size: Tensor,
391    /// Stride `(H, W)`.
392    pub stride: Tensor,
393    /// Padding `(H, W)`.
394    pub padding: Tensor,
395    /// Free-form node metadata (Python `metadata` dict).
396    pub metadata: MetadataMap,
397}
398
399/// Heaviside threshold / surrogate step.
400#[derive(Debug, Clone, PartialEq)]
401#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
402pub struct Threshold {
403    /// Threshold value(s).
404    pub threshold: Tensor,
405    /// Free-form node metadata (Python `metadata` dict).
406    pub metadata: MetadataMap,
407}
408
409#[cfg(test)]
410mod tests {
411    use super::*;
412    use crate::types::Tensor;
413
414    fn sample_weight() -> Tensor {
415        Tensor::from_f32(vec![2, 3], vec![1., 0., 0., 0., 1., 0.]).unwrap()
416    }
417
418    fn sample_bias() -> Tensor {
419        Tensor::from_f32(vec![2], vec![0., 0.]).unwrap()
420    }
421
422    fn sample_vec3() -> Tensor {
423        Tensor::from_f64(vec![3], vec![1.0, 1.0, 1.0]).unwrap()
424    }
425
426    /// Shared `(kernel_size, stride, padding)` tensors for SumPool2d / AvgPool2d
427    /// wire-name samples (keeps the two cases from being pure copy-paste).
428    fn sample_pool2d_fields() -> (Tensor, Tensor, Tensor) {
429        (
430            Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
431            Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
432            Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
433        )
434    }
435
436    #[test]
437    fn all_type_names_match_wire_strings() {
438        let cases: Vec<(&str, NirNode)> = vec![
439            (
440                "Input",
441                NirNode::Input(Input {
442                    shape: vec![1, 4],
443                    metadata: Default::default(),
444                }),
445            ),
446            (
447                "Output",
448                NirNode::Output(Output {
449                    shape: vec![1, 2],
450                    metadata: Default::default(),
451                }),
452            ),
453            (
454                "Affine",
455                NirNode::Affine(Affine {
456                    weight: sample_weight(),
457                    bias: sample_bias(),
458                    metadata: Default::default(),
459                }),
460            ),
461            (
462                "Linear",
463                NirNode::Linear(Linear {
464                    weight: sample_weight(),
465                    metadata: Default::default(),
466                }),
467            ),
468            (
469                "Scale",
470                NirNode::Scale(Scale {
471                    scale: sample_vec3(),
472                    metadata: Default::default(),
473                }),
474            ),
475            (
476                "Conv1d",
477                NirNode::Conv1d(Conv1d {
478                    weight: Tensor::from_f32(vec![1, 1, 3], vec![1., 0., -1.]).unwrap(),
479                    stride: vec![1],
480                    padding: Padding::single(0),
481                    dilation: vec![1],
482                    groups: 1,
483                    bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
484                    input_shape: Some(10),
485                    metadata: Default::default(),
486                }),
487            ),
488            (
489                "Conv2d",
490                NirNode::Conv2d(Conv2d {
491                    weight: Tensor::from_f32(vec![1, 1, 3, 3], vec![0.; 9]).unwrap(),
492                    stride: vec![1, 1],
493                    padding: Padding::Same,
494                    dilation: vec![1, 1],
495                    groups: 1,
496                    bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
497                    input_shape: Some(vec![28, 28]),
498                    metadata: Default::default(),
499                }),
500            ),
501            (
502                "CubaLI",
503                NirNode::CubaLi(CubaLi {
504                    tau_syn: sample_vec3(),
505                    tau_mem: sample_vec3(),
506                    r: sample_vec3(),
507                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
508                    w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
509                    metadata: Default::default(),
510                }),
511            ),
512            (
513                "CubaLIF",
514                NirNode::CubaLif(CubaLif {
515                    tau_syn: sample_vec3(),
516                    tau_mem: sample_vec3(),
517                    r: sample_vec3(),
518                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
519                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
520                    v_reset: None,
521                    w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
522                    metadata: Default::default(),
523                }),
524            ),
525            (
526                "Delay",
527                NirNode::Delay(Delay {
528                    delay: Tensor::scalar_f64(1.0),
529                    metadata: Default::default(),
530                }),
531            ),
532            (
533                "Flatten",
534                NirNode::Flatten(Flatten {
535                    start_dim: 1,
536                    end_dim: -1,
537                    input_type: Some(vec![1, 4, 4]),
538                    metadata: Default::default(),
539                }),
540            ),
541            (
542                "I",
543                NirNode::I(I {
544                    r: sample_vec3(),
545                    metadata: Default::default(),
546                }),
547            ),
548            (
549                "IF",
550                NirNode::If(If {
551                    r: sample_vec3(),
552                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
553                    v_reset: None,
554                    metadata: Default::default(),
555                }),
556            ),
557            (
558                "LI",
559                NirNode::Li(Li {
560                    tau: sample_vec3(),
561                    r: sample_vec3(),
562                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
563                    metadata: Default::default(),
564                }),
565            ),
566            (
567                "LIF",
568                NirNode::Lif(Lif {
569                    tau: sample_vec3(),
570                    r: sample_vec3(),
571                    v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
572                    v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
573                    v_reset: Some(Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap()),
574                    metadata: Default::default(),
575                }),
576            ),
577            {
578                let (kernel_size, stride, padding) = sample_pool2d_fields();
579                (
580                    "SumPool2d",
581                    NirNode::SumPool2d(SumPool2d {
582                        kernel_size,
583                        stride,
584                        padding,
585                        metadata: Default::default(),
586                    }),
587                )
588            },
589            {
590                let (kernel_size, stride, padding) = sample_pool2d_fields();
591                (
592                    "AvgPool2d",
593                    NirNode::AvgPool2d(AvgPool2d {
594                        kernel_size,
595                        stride,
596                        padding,
597                        metadata: Default::default(),
598                    }),
599                )
600            },
601            (
602                "Threshold",
603                NirNode::Threshold(Threshold {
604                    threshold: Tensor::scalar_f64(1.0),
605                    metadata: Default::default(),
606                }),
607            ),
608            ("NIRGraph", NirNode::Graph(Box::new(NirGraph::new()))),
609        ];
610
611        assert_eq!(cases.len(), 19, "expected all wire node types");
612        for (wire, node) in cases {
613            assert_eq!(node.type_name(), wire);
614            #[cfg(feature = "serde")]
615            {
616                let value = serde_json::to_value(&node).unwrap();
617                assert_eq!(value["type"], wire);
618                assert_eq!(serde_json::from_value::<NirNode>(value).unwrap(), node);
619            }
620        }
621    }
622
623    #[test]
624    fn padding_helpers() {
625        assert_eq!(Padding::single(1), Padding::Explicit(vec![1]));
626        assert_eq!(Padding::pair(1, 2), Padding::Explicit(vec![1, 2]));
627        let _ = Padding::Same;
628        let _ = Padding::Valid;
629    }
630
631    #[test]
632    fn never_use_marketing_aliases() {
633        // Guard against accidental CurrLIF / Convolution / Integrator names.
634        let names: Vec<&str> = [
635            NirNode::CubaLif(CubaLif {
636                tau_syn: Tensor::scalar_f64(1.0),
637                tau_mem: Tensor::scalar_f64(1.0),
638                r: Tensor::scalar_f64(1.0),
639                v_leak: Tensor::scalar_f64(0.0),
640                v_threshold: Tensor::scalar_f64(1.0),
641                v_reset: None,
642                w_in: Some(Tensor::scalar_f64(1.0)),
643                metadata: Default::default(),
644            }),
645            NirNode::Conv2d(Conv2d {
646                weight: Tensor::from_f32(vec![1, 1, 1, 1], vec![1.]).unwrap(),
647                stride: vec![1, 1],
648                padding: Padding::Valid,
649                dilation: vec![1, 1],
650                groups: 1,
651                bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
652                input_shape: None,
653                metadata: Default::default(),
654            }),
655            NirNode::I(I {
656                r: Tensor::scalar_f64(1.0),
657                metadata: Default::default(),
658            }),
659            NirNode::SumPool2d(SumPool2d {
660                kernel_size: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
661                stride: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
662                padding: Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
663                metadata: Default::default(),
664            }),
665        ]
666        .into_iter()
667        .map(|n| n.type_name())
668        .collect();
669
670        assert_eq!(names, ["CubaLIF", "Conv2d", "I", "SumPool2d"]);
671        for n in names {
672            assert!(!n.contains("Curr"));
673            assert!(!n.contains("Convolution"));
674            assert!(!n.contains("Integrator"));
675            assert!(!n.contains("Pooling"));
676        }
677    }
678}