use crate::error::{NirError, Result};
use crate::nodes::Padding;
pub const KEY_VERSION: &str = "version";
pub const KEY_NODE: &str = "node";
pub const KEY_NODES: &str = "nodes";
pub const KEY_EDGES: &str = "edges";
pub const KEY_METADATA: &str = "metadata";
pub const KEY_TYPE: &str = "type";
pub const WIRE_TYPES: [&str; 19] = [
"Input",
"Output",
"Affine",
"Linear",
"Scale",
"Conv1d",
"Conv2d",
"CubaLI",
"CubaLIF",
"Delay",
"Flatten",
"I",
"IF",
"LI",
"LIF",
"SumPool2d",
"AvgPool2d",
"Threshold",
"NIRGraph",
];
#[must_use]
pub fn is_wire_type(name: &str) -> bool {
WIRE_TYPES.contains(&name)
}
#[must_use]
pub fn padding_as_wire_str(padding: &Padding) -> Option<&'static str> {
match padding {
Padding::Same => Some("same"),
Padding::Valid => Some("valid"),
Padding::Explicit(_) => None,
}
}
pub fn padding_from_wire_str(s: &str) -> Result<Padding> {
match s {
"same" => Ok(Padding::Same),
"valid" => Ok(Padding::Valid),
other => Err(NirError::InvalidGraph(format!(
"padding must be \"same\", \"valid\", or integer extents, not {other:?}"
))),
}
}
pub fn check_link_name(kind: &str, name: &str) -> Result<()> {
if name.is_empty() {
return Err(NirError::InvalidGraph(format!("{kind} must not be empty")));
}
if name.contains('/') {
return Err(NirError::InvalidGraph(format!(
"{kind} {name:?} must not contain '/' (HDF5 path separator)"
)));
}
if name.contains('\0') {
return Err(NirError::InvalidGraph(format!(
"{kind} {name:?} must not contain a NUL byte (HDF5 link names are C strings)"
)));
}
if name == "." || name == ".." {
return Err(NirError::InvalidGraph(format!(
"{kind} {name:?} is a reserved HDF5 path component"
)));
}
Ok(())
}
pub fn check_hdf5_string(kind: &str, value: &str) -> Result<()> {
if value.contains('\0') {
return Err(NirError::InvalidGraph(format!(
"{kind} {value:?} must not contain a NUL byte (HDF5 strings are C strings)"
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::NirGraph;
use crate::nodes::{
Affine, AvgPool2d, Conv1d, Conv2d, CubaLi, CubaLif, Delay, Flatten, I, If, Input, Li, Lif,
Linear, NirNode, Output, Scale, SumPool2d, Threshold,
};
use crate::types::Tensor;
fn v() -> Tensor {
Tensor::from_f64([2], vec![1.0, 1.0]).unwrap()
}
fn pool() -> Tensor {
Tensor::from_i64([2], vec![2, 2]).unwrap()
}
fn port_and_linear_nodes() -> Vec<NirNode> {
let weight = || Tensor::from_f32(vec![2, 2], vec![1., 0., 0., 1.]).unwrap();
vec![
NirNode::Input(Input {
shape: vec![2],
metadata: Default::default(),
}),
NirNode::Output(Output {
shape: vec![2],
metadata: Default::default(),
}),
NirNode::Affine(Affine {
weight: weight(),
bias: Tensor::from_f32([2], vec![0., 0.]).unwrap(),
metadata: Default::default(),
}),
NirNode::Linear(Linear {
weight: weight(),
metadata: Default::default(),
}),
NirNode::Scale(Scale {
scale: v(),
metadata: Default::default(),
}),
]
}
fn conv_nodes() -> Vec<NirNode> {
vec![
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([1], vec![0.]).unwrap(),
input_shape: Some(10),
metadata: Default::default(),
}),
NirNode::Conv2d(Conv2d {
weight: Tensor::from_f32(vec![1, 1, 2, 2], vec![0.; 4]).unwrap(),
stride: vec![1, 1],
padding: Padding::Same,
dilation: vec![1, 1],
groups: 1,
bias: Tensor::from_f32([1], vec![0.]).unwrap(),
input_shape: Some(vec![8, 8]),
metadata: Default::default(),
}),
]
}
fn cuba_nodes() -> Vec<NirNode> {
vec![
NirNode::CubaLi(CubaLi {
tau_syn: v(),
tau_mem: v(),
r: v(),
v_leak: v(),
w_in: None,
metadata: Default::default(),
}),
NirNode::CubaLif(CubaLif {
tau_syn: v(),
tau_mem: v(),
r: v(),
v_leak: v(),
v_threshold: v(),
v_reset: None,
w_in: None,
metadata: Default::default(),
}),
]
}
fn neuron_nodes() -> Vec<NirNode> {
vec![
NirNode::I(I {
r: v(),
metadata: Default::default(),
}),
NirNode::If(If {
r: v(),
v_threshold: v(),
v_reset: None,
metadata: Default::default(),
}),
NirNode::Li(Li {
tau: v(),
r: v(),
v_leak: v(),
metadata: Default::default(),
}),
NirNode::Lif(Lif {
tau: v(),
r: v(),
v_leak: v(),
v_threshold: v(),
v_reset: None,
metadata: Default::default(),
}),
]
}
fn pool_nodes() -> Vec<NirNode> {
let no_pad = || Tensor::from_i64([2], vec![0, 0]).unwrap();
vec![
NirNode::SumPool2d(SumPool2d {
kernel_size: pool(),
stride: pool(),
padding: no_pad(),
metadata: Default::default(),
}),
NirNode::AvgPool2d(AvgPool2d {
kernel_size: pool(),
stride: pool(),
padding: no_pad(),
metadata: Default::default(),
}),
]
}
fn variant_index(node: &NirNode) -> usize {
match node {
NirNode::Input(_) => 0,
NirNode::Output(_) => 1,
NirNode::Affine(_) => 2,
NirNode::Linear(_) => 3,
NirNode::Scale(_) => 4,
NirNode::Conv1d(_) => 5,
NirNode::Conv2d(_) => 6,
NirNode::CubaLi(_) => 7,
NirNode::CubaLif(_) => 8,
NirNode::Delay(_) => 9,
NirNode::Flatten(_) => 10,
NirNode::I(_) => 11,
NirNode::If(_) => 12,
NirNode::Li(_) => 13,
NirNode::Lif(_) => 14,
NirNode::SumPool2d(_) => 15,
NirNode::AvgPool2d(_) => 16,
NirNode::Threshold(_) => 17,
NirNode::Graph(_) => 18,
}
}
fn one_of_each() -> Vec<NirNode> {
let mut nodes = port_and_linear_nodes();
nodes.extend(conv_nodes());
nodes.extend(cuba_nodes());
nodes.push(NirNode::Delay(Delay {
delay: v(),
metadata: Default::default(),
}));
nodes.push(NirNode::Flatten(Flatten {
start_dim: 1,
end_dim: -1,
input_type: None,
metadata: Default::default(),
}));
nodes.extend(neuron_nodes());
nodes.extend(pool_nodes());
nodes.push(NirNode::Threshold(Threshold {
threshold: v(),
metadata: Default::default(),
}));
nodes.push(NirNode::Graph(Box::new(NirGraph::new())));
nodes
}
#[test]
fn wire_types_matches_every_node_variant() {
const VARIANT_COUNT: usize = 19;
assert_eq!(WIRE_TYPES.len(), VARIANT_COUNT);
let nodes = one_of_each();
assert_eq!(nodes.len(), VARIANT_COUNT);
let mut seen = [false; VARIANT_COUNT];
for node in &nodes {
let i = variant_index(node);
assert!(
i < VARIANT_COUNT,
"variant_index {i} is outside VARIANT_COUNT"
);
assert!(!seen[i], "duplicate sample for variant index {i}");
seen[i] = true;
assert_eq!(
node.type_name(),
WIRE_TYPES[i],
"sample at index {i} must match WIRE_TYPES"
);
}
assert!(
seen.iter().all(|&s| s),
"one_of_each must cover every variant index"
);
}
#[test]
fn is_wire_type_rejects_marketing_aliases() {
for good in WIRE_TYPES {
assert!(is_wire_type(good), "{good} should be a wire type");
}
for bad in ["CurrLIF", "Convolution", "Integrator", "SumPooling", ""] {
assert!(!is_wire_type(bad), "{bad} must not be a wire type");
}
}
#[test]
fn padding_wire_strings_round_trip() {
assert_eq!(padding_as_wire_str(&Padding::Same), Some("same"));
assert_eq!(padding_as_wire_str(&Padding::Valid), Some("valid"));
assert_eq!(padding_as_wire_str(&Padding::pair(1, 1)), None);
assert_eq!(padding_from_wire_str("same").unwrap(), Padding::Same);
assert_eq!(padding_from_wire_str("valid").unwrap(), Padding::Valid);
}
#[test]
fn padding_from_unknown_string_is_rejected() {
let err = padding_from_wire_str("SAME").unwrap_err();
assert!(matches!(err, NirError::InvalidGraph(_)));
assert!(err.to_string().contains("\"SAME\""));
}
#[test]
fn node_names_with_dots_are_allowed() {
assert!(check_link_name("node name", "lif1.lif").is_ok());
assert!(check_link_name("node name", "0").is_ok());
}
#[test]
fn illegal_link_names_are_rejected() {
for bad in ["", "a/b", ".", "..", "nul\0inside"] {
let err = check_link_name("node name", bad).unwrap_err();
assert!(
matches!(err, NirError::InvalidGraph(_)),
"{bad:?} should be rejected"
);
}
}
#[test]
fn the_kind_label_appears_in_the_error() {
let err = check_link_name("metadata key", "a/b").unwrap_err();
assert!(err.to_string().contains("metadata key"), "got {err}");
}
#[test]
fn nul_bytes_in_string_values_are_rejected() {
let err = check_hdf5_string("version", "0.2\0.0").unwrap_err();
assert!(matches!(err, NirError::InvalidGraph(_)));
assert!(err.to_string().contains("version"), "got {err}");
assert!(err.to_string().contains("NUL"), "got {err}");
}
}