use prost::Message;
#[derive(Clone, PartialEq, Message)]
struct ModelProto {
#[prost(int64, tag = "1")]
ir_version: i64,
#[prost(message, repeated, tag = "8")]
opset_import: Vec<OperatorSetIdProto>,
#[prost(message, optional, tag = "7")]
graph: Option<GraphProto>,
#[prost(message, repeated, tag = "25")]
functions: Vec<FunctionProto>,
}
#[derive(Clone, PartialEq, Message)]
struct OperatorSetIdProto {
#[prost(string, tag = "1")]
domain: String,
#[prost(int64, tag = "2")]
version: i64,
}
#[derive(Clone, PartialEq, Message)]
struct GraphProto {
#[prost(message, repeated, tag = "1")]
node: Vec<NodeProto>,
#[prost(message, repeated, tag = "5")]
initializer: Vec<TensorProto>,
#[prost(message, repeated, tag = "15")]
sparse_initializer: Vec<SparseTensorProto>,
#[prost(message, repeated, tag = "11")]
input: Vec<ValueInfoProto>,
#[prost(message, repeated, tag = "12")]
output: Vec<ValueInfoProto>,
}
#[derive(Clone, PartialEq, Message)]
struct FunctionProto {
#[prost(message, repeated, tag = "7")]
node: Vec<NodeProto>,
}
#[derive(Clone, PartialEq, Message)]
struct NodeProto {
#[prost(string, tag = "4")]
op_type: String,
#[prost(string, tag = "7")]
domain: String,
#[prost(message, repeated, tag = "5")]
attribute: Vec<AttributeProto>,
}
impl NodeProto {
fn operator(&self) -> String {
if self.domain.is_empty() || self.domain == "ai.onnx" {
self.op_type.clone()
} else {
format!("{}.{}", self.domain, self.op_type)
}
}
}
#[derive(Clone, PartialEq, Message)]
struct AttributeProto {
#[prost(message, optional, tag = "5")]
t: Option<TensorProto>,
#[prost(message, optional, boxed, tag = "6")]
g: Option<Box<GraphProto>>,
#[prost(float, repeated, tag = "7")]
floats: Vec<f32>,
#[prost(int64, repeated, tag = "8")]
ints: Vec<i64>,
#[prost(message, repeated, tag = "10")]
tensors: Vec<TensorProto>,
#[prost(message, repeated, tag = "11")]
graphs: Vec<GraphProto>,
#[prost(message, optional, tag = "22")]
sparse_tensor: Option<SparseTensorProto>,
#[prost(message, repeated, tag = "23")]
sparse_tensors: Vec<SparseTensorProto>,
}
#[derive(Clone, PartialEq, Message)]
struct TensorProto {
#[prost(int64, repeated, tag = "1")]
dims: Vec<i64>,
#[prost(string, tag = "8")]
name: String,
}
#[derive(Clone, PartialEq, Message)]
struct SparseTensorProto {
#[prost(message, optional, tag = "1")]
values: Option<TensorProto>,
}
#[derive(Clone, PartialEq, Message)]
struct ValueInfoProto {
#[prost(string, tag = "1")]
name: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GraphStats {
pub parameters: u64,
pub nodes: u64,
pub operators: Vec<String>,
pub ir_version: i64,
pub opset: i64,
pub input_names: Vec<String>,
pub output_names: Vec<String>,
}
fn tensor_values(tensor: &TensorProto) -> u64 {
tensor.dims.iter().fold(1u64, |n, d| {
n.saturating_mul(u64::try_from(*d).unwrap_or(0))
})
}
fn sparse_values(sparse: &SparseTensorProto) -> u64 {
sparse.values.as_ref().map_or(0, tensor_values)
}
fn stored_values(graph: &GraphProto) -> u64 {
let dense = graph.initializer.iter().fold(0u64, |sum, tensor| {
sum.saturating_add(tensor_values(tensor))
});
graph.sparse_initializer.iter().fold(dense, |sum, sparse| {
sum.saturating_add(sparse_values(sparse))
})
}
fn carried_values(attribute: &AttributeProto) -> u64 {
let mut values = attribute.t.as_ref().map_or(0, tensor_values);
for tensor in &attribute.tensors {
values = values.saturating_add(tensor_values(tensor));
}
for sparse in attribute
.sparse_tensor
.iter()
.chain(&attribute.sparse_tensors)
{
values = values.saturating_add(sparse_values(sparse));
}
values
.saturating_add(attribute.floats.len() as u64)
.saturating_add(attribute.ints.len() as u64)
}
pub fn read_stats(bytes: &[u8]) -> Result<GraphStats, String> {
let model = ModelProto::decode(bytes).map_err(|e| format!("not an ONNX model: {e}"))?;
let Some(graph) = model.graph.as_ref() else {
return Err("not an ONNX model: the document carries no graph".to_string());
};
if model.ir_version <= 0 {
return Err("not an ONNX model: the document declares no IR version".to_string());
}
let mut parameters = stored_values(graph);
let mut nodes = 0u64;
let mut operators = std::collections::BTreeSet::new();
let mut pending: Vec<&[NodeProto]> = vec![&graph.node];
pending.extend(model.functions.iter().map(|f| f.node.as_slice()));
while let Some(body) = pending.pop() {
nodes = nodes.saturating_add(body.len() as u64);
for node in body {
operators.insert(node.operator());
for attribute in &node.attribute {
parameters = parameters.saturating_add(carried_values(attribute));
for sub in attribute.g.as_deref().into_iter().chain(&attribute.graphs) {
parameters = parameters.saturating_add(stored_values(sub));
pending.push(&sub.node);
}
}
}
}
let opset = model
.opset_import
.iter()
.find(|o| o.domain.is_empty() || o.domain == "ai.onnx")
.map_or(0, |o| o.version);
let initialized = |name: &str| {
graph.initializer.iter().any(|t| t.name == name)
|| graph
.sparse_initializer
.iter()
.any(|s| s.values.as_ref().is_some_and(|t| t.name == name))
};
let input_names = graph
.input
.iter()
.filter(|input| !initialized(&input.name))
.map(|input| input.name.clone())
.collect();
let output_names = graph.output.iter().map(|o| o.name.clone()).collect();
Ok(GraphStats {
parameters,
nodes,
operators: operators.into_iter().collect(),
ir_version: model.ir_version,
opset,
input_names,
output_names,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::fixture;
#[test]
fn the_fixture_reads_back_as_built() {
let stats = read_stats(fixture::ONNX).expect("the fixture is an ONNX model");
assert_eq!(
stats,
GraphStats {
parameters: 1479,
nodes: 4,
operators: ["Flatten", "Gemm", "Relu"].map(str::to_string).to_vec(),
ir_version: 9,
opset: 17,
input_names: vec!["board".to_string()],
output_names: vec!["policy".to_string()],
}
);
}
#[test]
fn the_same_network_counts_the_same_however_its_weights_are_carried() {
let as_init = read_stats(fixture::AS_INIT_ONNX).expect("an ONNX model");
let as_const = read_stats(fixture::AS_CONST_ONNX).expect("an ONNX model");
let as_list = read_stats(fixture::AS_LIST_ONNX).expect("an ONNX model");
assert_eq!(as_init.parameters, 15);
assert_eq!(as_const.parameters, 15);
assert_eq!(as_list.parameters, 17);
assert_eq!((as_init.nodes, as_const.nodes, as_list.nodes), (1, 3, 5));
assert_eq!(as_init.operators, ["Gemm"]);
assert_eq!(as_const.operators, ["Constant", "Gemm"]);
assert_eq!(as_list.operators, ["Constant", "Gemm", "Reshape"]);
assert_eq!(as_init.input_names, as_const.input_names);
assert_eq!(as_init.output_names, as_list.output_names);
}
#[test]
fn initializers_are_parameters_and_not_inputs() {
let model = ModelProto {
ir_version: 7,
opset_import: vec![
OperatorSetIdProto {
domain: "com.microsoft".to_string(),
version: 1,
},
OperatorSetIdProto {
domain: String::new(),
version: 13,
},
],
graph: Some(GraphProto {
node: vec![NodeProto::default(), NodeProto::default()],
initializer: vec![
TensorProto {
dims: vec![3, 4],
name: "w".to_string(),
},
TensorProto {
dims: vec![],
name: "scale".to_string(),
},
],
input: vec![
ValueInfoProto {
name: "x".to_string(),
},
ValueInfoProto {
name: "w".to_string(),
},
],
output: vec![ValueInfoProto {
name: "y".to_string(),
}],
..GraphProto::default()
}),
functions: vec![],
};
let stats = read_stats(&model.encode_to_vec()).expect("decodes");
assert_eq!(stats.parameters, 13);
assert_eq!(stats.nodes, 2);
assert_eq!(stats.ir_version, 7);
assert_eq!(stats.opset, 13);
assert_eq!(stats.input_names, ["x"]);
assert_eq!(stats.output_names, ["y"]);
}
#[test]
fn a_value_counts_wherever_the_document_carries_it() {
let tensor = |dims: Vec<i64>| TensorProto {
dims,
name: String::new(),
};
let holding = |op_type: &str, attribute: AttributeProto| NodeProto {
op_type: op_type.to_string(),
domain: String::new(),
attribute: vec![attribute],
};
let inner = GraphProto {
initializer: vec![tensor(vec![4])],
..GraphProto::default()
};
let model = ModelProto {
ir_version: 9,
opset_import: vec![OperatorSetIdProto {
domain: String::new(),
version: 17,
}],
graph: Some(GraphProto {
node: vec![
holding(
"Constant",
AttributeProto {
t: Some(tensor(vec![2, 3])),
..AttributeProto::default()
},
),
holding(
"Constant",
AttributeProto {
floats: vec![0.0; 6],
..AttributeProto::default()
},
),
holding(
"Conv",
AttributeProto {
ints: vec![3, 3],
..AttributeProto::default()
},
),
holding(
"If",
AttributeProto {
g: Some(Box::new(GraphProto {
node: vec![holding(
"Loop",
AttributeProto {
g: Some(Box::new(inner)),
..AttributeProto::default()
},
)],
initializer: vec![tensor(vec![5])],
..GraphProto::default()
})),
..AttributeProto::default()
},
),
],
sparse_initializer: vec![SparseTensorProto {
values: Some(tensor(vec![2])),
}],
..GraphProto::default()
}),
functions: vec![FunctionProto {
node: vec![NodeProto {
op_type: "LinearRegressor".to_string(),
domain: "ai.onnx.ml".to_string(),
attribute: vec![AttributeProto {
t: Some(tensor(vec![7])),
..AttributeProto::default()
}],
}],
}],
};
let stats = read_stats(&model.encode_to_vec()).expect("decodes");
assert_eq!(stats.parameters, 32);
assert_eq!(stats.nodes, 6);
assert_eq!(
stats.operators,
[
"Constant",
"Conv",
"If",
"Loop",
"ai.onnx.ml.LinearRegressor"
]
);
}
#[test]
fn what_is_not_a_model_says_so() {
let err = read_stats(b"\x00\x01not a protobuf at all").expect_err("garbage");
assert!(err.starts_with("not an ONNX model:"), "{err}");
let err = read_stats(b"").expect_err("empty");
assert!(err.contains("no graph"), "{err}");
let graph_only = ModelProto {
ir_version: 0,
graph: Some(GraphProto::default()),
..ModelProto::default()
};
let err = read_stats(&graph_only.encode_to_vec()).expect_err("no ir_version");
assert!(err.contains("IR version"), "{err}");
}
}