1use crate::graph::NirGraph;
14use crate::types::{MetadataMap, Tensor};
15
16#[derive(Debug, Clone, PartialEq)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21#[non_exhaustive]
22pub enum Padding {
23 Explicit(Vec<i64>),
25 Same,
27 Valid,
29}
30
31impl Padding {
32 #[must_use]
34 pub fn single(value: i64) -> Self {
35 Self::Explicit(vec![value])
36 }
37
38 #[must_use]
40 pub fn pair(h: i64, w: i64) -> Self {
41 Self::Explicit(vec![h, w])
42 }
43}
44
45#[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 #[cfg_attr(feature = "serde", serde(rename = "Input"))]
58 Input(Input),
59 #[cfg_attr(feature = "serde", serde(rename = "Output"))]
61 Output(Output),
62 #[cfg_attr(feature = "serde", serde(rename = "Affine"))]
64 Affine(Affine),
65 #[cfg_attr(feature = "serde", serde(rename = "Linear"))]
67 Linear(Linear),
68 #[cfg_attr(feature = "serde", serde(rename = "Scale"))]
70 Scale(Scale),
71 #[cfg_attr(feature = "serde", serde(rename = "Conv1d"))]
73 Conv1d(Conv1d),
74 #[cfg_attr(feature = "serde", serde(rename = "Conv2d"))]
76 Conv2d(Conv2d),
77 #[cfg_attr(feature = "serde", serde(rename = "CubaLI"))]
79 CubaLi(CubaLi),
80 #[cfg_attr(feature = "serde", serde(rename = "CubaLIF"))]
82 CubaLif(CubaLif),
83 #[cfg_attr(feature = "serde", serde(rename = "Delay"))]
85 Delay(Delay),
86 #[cfg_attr(feature = "serde", serde(rename = "Flatten"))]
88 Flatten(Flatten),
89 #[cfg_attr(feature = "serde", serde(rename = "I"))]
91 I(I),
92 #[cfg_attr(feature = "serde", serde(rename = "IF"))]
94 If(If),
95 #[cfg_attr(feature = "serde", serde(rename = "LI"))]
97 Li(Li),
98 #[cfg_attr(feature = "serde", serde(rename = "LIF"))]
100 Lif(Lif),
101 #[cfg_attr(feature = "serde", serde(rename = "SumPool2d"))]
103 SumPool2d(SumPool2d),
104 #[cfg_attr(feature = "serde", serde(rename = "AvgPool2d"))]
106 AvgPool2d(AvgPool2d),
107 #[cfg_attr(feature = "serde", serde(rename = "Threshold"))]
109 Threshold(Threshold),
110 #[cfg_attr(feature = "serde", serde(rename = "NIRGraph"))]
112 Graph(Box<NirGraph>),
113}
114
115impl NirNode {
116 #[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 pub fn validate_parameters(&self) -> crate::error::Result<()> {
158 crate::validation::validate_node(self, crate::validation::ANONYMOUS_NODE)
159 }
160}
161
162#[derive(Debug, Clone, PartialEq)]
166#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
167pub struct Input {
168 pub shape: Vec<usize>,
170 pub metadata: MetadataMap,
172}
173
174#[derive(Debug, Clone, PartialEq)]
178#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
179pub struct Output {
180 pub shape: Vec<usize>,
182 pub metadata: MetadataMap,
184}
185
186#[derive(Debug, Clone, PartialEq)]
188#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
189pub struct Affine {
190 pub weight: Tensor,
192 pub bias: Tensor,
194 pub metadata: MetadataMap,
196}
197
198#[derive(Debug, Clone, PartialEq)]
200#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
201pub struct Linear {
202 pub weight: Tensor,
204 pub metadata: MetadataMap,
206}
207
208#[derive(Debug, Clone, PartialEq)]
210#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
211pub struct Scale {
212 pub scale: Tensor,
214 pub metadata: MetadataMap,
216}
217
218#[derive(Debug, Clone, PartialEq)]
220#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
221pub struct Conv1d {
222 pub weight: Tensor,
224 pub stride: Vec<i64>,
226 pub padding: Padding,
228 pub dilation: Vec<i64>,
230 pub groups: i64,
232 pub bias: Tensor,
234 pub input_shape: Option<usize>,
236 pub metadata: MetadataMap,
238}
239
240#[derive(Debug, Clone, PartialEq)]
242#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
243pub struct Conv2d {
244 pub weight: Tensor,
246 pub stride: Vec<i64>,
248 pub padding: Padding,
250 pub dilation: Vec<i64>,
252 pub groups: i64,
254 pub bias: Tensor,
256 pub input_shape: Option<Vec<usize>>,
258 pub metadata: MetadataMap,
260}
261
262#[derive(Debug, Clone, PartialEq)]
264#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
265pub struct CubaLi {
266 pub tau_syn: Tensor,
268 pub tau_mem: Tensor,
270 pub r: Tensor,
272 pub v_leak: Tensor,
274 pub w_in: Option<Tensor>,
280 pub metadata: MetadataMap,
282}
283
284#[derive(Debug, Clone, PartialEq)]
286#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
287pub struct CubaLif {
288 pub tau_syn: Tensor,
290 pub tau_mem: Tensor,
292 pub r: Tensor,
294 pub v_leak: Tensor,
296 pub v_threshold: Tensor,
298 pub v_reset: Option<Tensor>,
300 pub w_in: Option<Tensor>,
306 pub metadata: MetadataMap,
308}
309
310#[derive(Debug, Clone, PartialEq)]
312#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
313pub struct Delay {
314 pub delay: Tensor,
316 pub metadata: MetadataMap,
318}
319
320#[derive(Debug, Clone, PartialEq)]
322#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
323pub struct Flatten {
324 pub start_dim: i64,
326 pub end_dim: i64,
328 pub input_type: Option<Vec<usize>>,
330 pub metadata: MetadataMap,
332}
333
334#[derive(Debug, Clone, PartialEq)]
336#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
337pub struct I {
338 pub r: Tensor,
340 pub metadata: MetadataMap,
342}
343
344#[derive(Debug, Clone, PartialEq)]
346#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
347pub struct If {
348 pub r: Tensor,
350 pub v_threshold: Tensor,
352 pub v_reset: Option<Tensor>,
354 pub metadata: MetadataMap,
356}
357
358#[derive(Debug, Clone, PartialEq)]
360#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
361pub struct Li {
362 pub tau: Tensor,
364 pub r: Tensor,
366 pub v_leak: Tensor,
368 pub metadata: MetadataMap,
370}
371
372#[derive(Debug, Clone, PartialEq)]
374#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
375pub struct Lif {
376 pub tau: Tensor,
378 pub r: Tensor,
380 pub v_leak: Tensor,
382 pub v_threshold: Tensor,
384 pub v_reset: Option<Tensor>,
386 pub metadata: MetadataMap,
388}
389
390#[derive(Debug, Clone, PartialEq)]
392#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
393pub struct SumPool2d {
394 pub kernel_size: Tensor,
396 pub stride: Tensor,
398 pub padding: Tensor,
400 pub metadata: MetadataMap,
402}
403
404#[derive(Debug, Clone, PartialEq)]
406#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
407pub struct AvgPool2d {
408 pub kernel_size: Tensor,
410 pub stride: Tensor,
412 pub padding: Tensor,
414 pub metadata: MetadataMap,
416}
417
418#[derive(Debug, Clone, PartialEq)]
420#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
421pub struct Threshold {
422 pub threshold: Tensor,
424 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 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 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}