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
143#[derive(Debug, Clone, PartialEq)]
147#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
148pub struct Input {
149 pub shape: Vec<usize>,
151 pub metadata: MetadataMap,
153}
154
155#[derive(Debug, Clone, PartialEq)]
159#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
160pub struct Output {
161 pub shape: Vec<usize>,
163 pub metadata: MetadataMap,
165}
166
167#[derive(Debug, Clone, PartialEq)]
169#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
170pub struct Affine {
171 pub weight: Tensor,
173 pub bias: Tensor,
175 pub metadata: MetadataMap,
177}
178
179#[derive(Debug, Clone, PartialEq)]
181#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
182pub struct Linear {
183 pub weight: Tensor,
185 pub metadata: MetadataMap,
187}
188
189#[derive(Debug, Clone, PartialEq)]
191#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
192pub struct Scale {
193 pub scale: Tensor,
195 pub metadata: MetadataMap,
197}
198
199#[derive(Debug, Clone, PartialEq)]
201#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
202pub struct Conv1d {
203 pub weight: Tensor,
205 pub stride: Vec<i64>,
207 pub padding: Padding,
209 pub dilation: Vec<i64>,
211 pub groups: i64,
213 pub bias: Tensor,
215 pub input_shape: Option<usize>,
217 pub metadata: MetadataMap,
219}
220
221#[derive(Debug, Clone, PartialEq)]
223#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
224pub struct Conv2d {
225 pub weight: Tensor,
227 pub stride: Vec<i64>,
229 pub padding: Padding,
231 pub dilation: Vec<i64>,
233 pub groups: i64,
235 pub bias: Tensor,
237 pub input_shape: Option<Vec<usize>>,
239 pub metadata: MetadataMap,
241}
242
243#[derive(Debug, Clone, PartialEq)]
245#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
246pub struct CubaLi {
247 pub tau_syn: Tensor,
249 pub tau_mem: Tensor,
251 pub r: Tensor,
253 pub v_leak: Tensor,
255 pub w_in: Option<Tensor>,
261 pub metadata: MetadataMap,
263}
264
265#[derive(Debug, Clone, PartialEq)]
267#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
268pub struct CubaLif {
269 pub tau_syn: Tensor,
271 pub tau_mem: Tensor,
273 pub r: Tensor,
275 pub v_leak: Tensor,
277 pub v_threshold: Tensor,
279 pub v_reset: Option<Tensor>,
281 pub w_in: Option<Tensor>,
287 pub metadata: MetadataMap,
289}
290
291#[derive(Debug, Clone, PartialEq)]
293#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
294pub struct Delay {
295 pub delay: Tensor,
297 pub metadata: MetadataMap,
299}
300
301#[derive(Debug, Clone, PartialEq)]
303#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
304pub struct Flatten {
305 pub start_dim: i64,
307 pub end_dim: i64,
309 pub input_type: Option<Vec<usize>>,
311 pub metadata: MetadataMap,
313}
314
315#[derive(Debug, Clone, PartialEq)]
317#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
318pub struct I {
319 pub r: Tensor,
321 pub metadata: MetadataMap,
323}
324
325#[derive(Debug, Clone, PartialEq)]
327#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
328pub struct If {
329 pub r: Tensor,
331 pub v_threshold: Tensor,
333 pub v_reset: Option<Tensor>,
335 pub metadata: MetadataMap,
337}
338
339#[derive(Debug, Clone, PartialEq)]
341#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
342pub struct Li {
343 pub tau: Tensor,
345 pub r: Tensor,
347 pub v_leak: Tensor,
349 pub metadata: MetadataMap,
351}
352
353#[derive(Debug, Clone, PartialEq)]
355#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
356pub struct Lif {
357 pub tau: Tensor,
359 pub r: Tensor,
361 pub v_leak: Tensor,
363 pub v_threshold: Tensor,
365 pub v_reset: Option<Tensor>,
367 pub metadata: MetadataMap,
369}
370
371#[derive(Debug, Clone, PartialEq)]
373#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
374pub struct SumPool2d {
375 pub kernel_size: Tensor,
377 pub stride: Tensor,
379 pub padding: Tensor,
381 pub metadata: MetadataMap,
383}
384
385#[derive(Debug, Clone, PartialEq)]
387#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
388pub struct AvgPool2d {
389 pub kernel_size: Tensor,
391 pub stride: Tensor,
393 pub padding: Tensor,
395 pub metadata: MetadataMap,
397}
398
399#[derive(Debug, Clone, PartialEq)]
401#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
402pub struct Threshold {
403 pub threshold: Tensor,
405 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 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 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}