use crate::graph::NirGraph;
use crate::types::{MetadataMap, Tensor};
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum Padding {
Explicit(Vec<i64>),
Same,
Valid,
}
impl Padding {
#[must_use]
pub fn single(value: i64) -> Self {
Self::Explicit(vec![value])
}
#[must_use]
pub fn pair(h: i64, w: i64) -> Self {
Self::Explicit(vec![h, w])
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(tag = "type"))]
pub enum NirNode {
#[cfg_attr(feature = "serde", serde(rename = "Input"))]
Input(Input),
#[cfg_attr(feature = "serde", serde(rename = "Output"))]
Output(Output),
#[cfg_attr(feature = "serde", serde(rename = "Affine"))]
Affine(Affine),
#[cfg_attr(feature = "serde", serde(rename = "Linear"))]
Linear(Linear),
#[cfg_attr(feature = "serde", serde(rename = "Scale"))]
Scale(Scale),
#[cfg_attr(feature = "serde", serde(rename = "Conv1d"))]
Conv1d(Conv1d),
#[cfg_attr(feature = "serde", serde(rename = "Conv2d"))]
Conv2d(Conv2d),
#[cfg_attr(feature = "serde", serde(rename = "CubaLI"))]
CubaLi(CubaLi),
#[cfg_attr(feature = "serde", serde(rename = "CubaLIF"))]
CubaLif(CubaLif),
#[cfg_attr(feature = "serde", serde(rename = "Delay"))]
Delay(Delay),
#[cfg_attr(feature = "serde", serde(rename = "Flatten"))]
Flatten(Flatten),
#[cfg_attr(feature = "serde", serde(rename = "I"))]
I(I),
#[cfg_attr(feature = "serde", serde(rename = "IF"))]
If(If),
#[cfg_attr(feature = "serde", serde(rename = "LI"))]
Li(Li),
#[cfg_attr(feature = "serde", serde(rename = "LIF"))]
Lif(Lif),
#[cfg_attr(feature = "serde", serde(rename = "SumPool2d"))]
SumPool2d(SumPool2d),
#[cfg_attr(feature = "serde", serde(rename = "AvgPool2d"))]
AvgPool2d(AvgPool2d),
#[cfg_attr(feature = "serde", serde(rename = "Threshold"))]
Threshold(Threshold),
#[cfg_attr(feature = "serde", serde(rename = "NIRGraph"))]
Graph(Box<NirGraph>),
}
impl NirNode {
#[must_use]
pub fn type_name(&self) -> &'static str {
match self {
Self::Input(_) => "Input",
Self::Output(_) => "Output",
Self::Affine(_) => "Affine",
Self::Linear(_) => "Linear",
Self::Scale(_) => "Scale",
Self::Conv1d(_) => "Conv1d",
Self::Conv2d(_) => "Conv2d",
Self::CubaLi(_) => "CubaLI",
Self::CubaLif(_) => "CubaLIF",
Self::Delay(_) => "Delay",
Self::Flatten(_) => "Flatten",
Self::I(_) => "I",
Self::If(_) => "IF",
Self::Li(_) => "LI",
Self::Lif(_) => "LIF",
Self::SumPool2d(_) => "SumPool2d",
Self::AvgPool2d(_) => "AvgPool2d",
Self::Threshold(_) => "Threshold",
Self::Graph(_) => "NIRGraph",
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Input {
pub shape: Vec<usize>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Output {
pub shape: Vec<usize>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Affine {
pub weight: Tensor,
pub bias: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Linear {
pub weight: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Scale {
pub scale: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Conv1d {
pub weight: Tensor,
pub stride: Vec<i64>,
pub padding: Padding,
pub dilation: Vec<i64>,
pub groups: i64,
pub bias: Tensor,
pub input_shape: Option<usize>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Conv2d {
pub weight: Tensor,
pub stride: Vec<i64>,
pub padding: Padding,
pub dilation: Vec<i64>,
pub groups: i64,
pub bias: Tensor,
pub input_shape: Option<Vec<usize>>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CubaLi {
pub tau_syn: Tensor,
pub tau_mem: Tensor,
pub r: Tensor,
pub v_leak: Tensor,
pub w_in: Option<Tensor>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CubaLif {
pub tau_syn: Tensor,
pub tau_mem: Tensor,
pub r: Tensor,
pub v_leak: Tensor,
pub v_threshold: Tensor,
pub v_reset: Option<Tensor>,
pub w_in: Option<Tensor>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Delay {
pub delay: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Flatten {
pub start_dim: i64,
pub end_dim: i64,
pub input_type: Option<Vec<usize>>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct I {
pub r: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct If {
pub r: Tensor,
pub v_threshold: Tensor,
pub v_reset: Option<Tensor>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Li {
pub tau: Tensor,
pub r: Tensor,
pub v_leak: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Lif {
pub tau: Tensor,
pub r: Tensor,
pub v_leak: Tensor,
pub v_threshold: Tensor,
pub v_reset: Option<Tensor>,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct SumPool2d {
pub kernel_size: Tensor,
pub stride: Tensor,
pub padding: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct AvgPool2d {
pub kernel_size: Tensor,
pub stride: Tensor,
pub padding: Tensor,
pub metadata: MetadataMap,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Threshold {
pub threshold: Tensor,
pub metadata: MetadataMap,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Tensor;
fn sample_weight() -> Tensor {
Tensor::from_f32(vec![2, 3], vec![1., 0., 0., 0., 1., 0.]).unwrap()
}
fn sample_bias() -> Tensor {
Tensor::from_f32(vec![2], vec![0., 0.]).unwrap()
}
fn sample_vec3() -> Tensor {
Tensor::from_f64(vec![3], vec![1.0, 1.0, 1.0]).unwrap()
}
fn sample_pool2d_fields() -> (Tensor, Tensor, Tensor) {
(
Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
)
}
#[test]
fn all_type_names_match_wire_strings() {
let cases: Vec<(&str, NirNode)> = vec![
(
"Input",
NirNode::Input(Input {
shape: vec![1, 4],
metadata: Default::default(),
}),
),
(
"Output",
NirNode::Output(Output {
shape: vec![1, 2],
metadata: Default::default(),
}),
),
(
"Affine",
NirNode::Affine(Affine {
weight: sample_weight(),
bias: sample_bias(),
metadata: Default::default(),
}),
),
(
"Linear",
NirNode::Linear(Linear {
weight: sample_weight(),
metadata: Default::default(),
}),
),
(
"Scale",
NirNode::Scale(Scale {
scale: sample_vec3(),
metadata: Default::default(),
}),
),
(
"Conv1d",
NirNode::Conv1d(Conv1d {
weight: Tensor::from_f32(vec![1, 1, 3], vec![1., 0., -1.]).unwrap(),
stride: vec![1],
padding: Padding::single(0),
dilation: vec![1],
groups: 1,
bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
input_shape: Some(10),
metadata: Default::default(),
}),
),
(
"Conv2d",
NirNode::Conv2d(Conv2d {
weight: Tensor::from_f32(vec![1, 1, 3, 3], vec![0.; 9]).unwrap(),
stride: vec![1, 1],
padding: Padding::Same,
dilation: vec![1, 1],
groups: 1,
bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
input_shape: Some(vec![28, 28]),
metadata: Default::default(),
}),
),
(
"CubaLI",
NirNode::CubaLi(CubaLi {
tau_syn: sample_vec3(),
tau_mem: sample_vec3(),
r: sample_vec3(),
v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
metadata: Default::default(),
}),
),
(
"CubaLIF",
NirNode::CubaLif(CubaLif {
tau_syn: sample_vec3(),
tau_mem: sample_vec3(),
r: sample_vec3(),
v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
v_reset: None,
w_in: Some(Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap()),
metadata: Default::default(),
}),
),
(
"Delay",
NirNode::Delay(Delay {
delay: Tensor::scalar_f64(1.0),
metadata: Default::default(),
}),
),
(
"Flatten",
NirNode::Flatten(Flatten {
start_dim: 1,
end_dim: -1,
input_type: Some(vec![1, 4, 4]),
metadata: Default::default(),
}),
),
(
"I",
NirNode::I(I {
r: sample_vec3(),
metadata: Default::default(),
}),
),
(
"IF",
NirNode::If(If {
r: sample_vec3(),
v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
v_reset: None,
metadata: Default::default(),
}),
),
(
"LI",
NirNode::Li(Li {
tau: sample_vec3(),
r: sample_vec3(),
v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
metadata: Default::default(),
}),
),
(
"LIF",
NirNode::Lif(Lif {
tau: sample_vec3(),
r: sample_vec3(),
v_leak: Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap(),
v_threshold: Tensor::from_f64(vec![3], vec![1., 1., 1.]).unwrap(),
v_reset: Some(Tensor::from_f64(vec![3], vec![0., 0., 0.]).unwrap()),
metadata: Default::default(),
}),
),
{
let (kernel_size, stride, padding) = sample_pool2d_fields();
(
"SumPool2d",
NirNode::SumPool2d(SumPool2d {
kernel_size,
stride,
padding,
metadata: Default::default(),
}),
)
},
{
let (kernel_size, stride, padding) = sample_pool2d_fields();
(
"AvgPool2d",
NirNode::AvgPool2d(AvgPool2d {
kernel_size,
stride,
padding,
metadata: Default::default(),
}),
)
},
(
"Threshold",
NirNode::Threshold(Threshold {
threshold: Tensor::scalar_f64(1.0),
metadata: Default::default(),
}),
),
("NIRGraph", NirNode::Graph(Box::new(NirGraph::new()))),
];
assert_eq!(cases.len(), 19, "expected all wire node types");
for (wire, node) in cases {
assert_eq!(node.type_name(), wire);
#[cfg(feature = "serde")]
{
let value = serde_json::to_value(&node).unwrap();
assert_eq!(value["type"], wire);
assert_eq!(serde_json::from_value::<NirNode>(value).unwrap(), node);
}
}
}
#[test]
fn padding_helpers() {
assert_eq!(Padding::single(1), Padding::Explicit(vec![1]));
assert_eq!(Padding::pair(1, 2), Padding::Explicit(vec![1, 2]));
let _ = Padding::Same;
let _ = Padding::Valid;
}
#[test]
fn never_use_marketing_aliases() {
let names: Vec<&str> = [
NirNode::CubaLif(CubaLif {
tau_syn: Tensor::scalar_f64(1.0),
tau_mem: Tensor::scalar_f64(1.0),
r: Tensor::scalar_f64(1.0),
v_leak: Tensor::scalar_f64(0.0),
v_threshold: Tensor::scalar_f64(1.0),
v_reset: None,
w_in: Some(Tensor::scalar_f64(1.0)),
metadata: Default::default(),
}),
NirNode::Conv2d(Conv2d {
weight: Tensor::from_f32(vec![1, 1, 1, 1], vec![1.]).unwrap(),
stride: vec![1, 1],
padding: Padding::Valid,
dilation: vec![1, 1],
groups: 1,
bias: Tensor::from_f32(vec![1], vec![0.]).unwrap(),
input_shape: None,
metadata: Default::default(),
}),
NirNode::I(I {
r: Tensor::scalar_f64(1.0),
metadata: Default::default(),
}),
NirNode::SumPool2d(SumPool2d {
kernel_size: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
stride: Tensor::from_i64(vec![2], vec![2, 2]).unwrap(),
padding: Tensor::from_i64(vec![2], vec![0, 0]).unwrap(),
metadata: Default::default(),
}),
]
.into_iter()
.map(|n| n.type_name())
.collect();
assert_eq!(names, ["CubaLIF", "Conv2d", "I", "SumPool2d"]);
for n in names {
assert!(!n.contains("Curr"));
assert!(!n.contains("Convolution"));
assert!(!n.contains("Integrator"));
assert!(!n.contains("Pooling"));
}
}
}