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>,
}
#[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 = "11")]
input: Vec<ValueInfoProto>,
#[prost(message, repeated, tag = "12")]
output: Vec<ValueInfoProto>,
}
#[derive(Clone, PartialEq, Message)]
struct NodeProto {
#[prost(string, repeated, tag = "2")]
output: Vec<String>,
#[prost(string, tag = "4")]
op_type: String,
}
#[derive(Clone, PartialEq, Message)]
struct TensorProto {
#[prost(int64, repeated, tag = "1")]
dims: Vec<i64>,
#[prost(int32, tag = "2")]
data_type: i32,
#[prost(string, tag = "8")]
name: String,
}
#[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 ir_version: i64,
pub opset: i64,
pub input_names: Vec<String>,
pub output_names: Vec<String>,
}
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 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 parameters = graph.initializer.iter().fold(0u64, |sum, tensor| {
let count = tensor.dims.iter().fold(1u64, |n, d| {
n.saturating_mul(u64::try_from(*d).unwrap_or(0))
});
sum.saturating_add(count)
});
let opset = model
.opset_import
.iter()
.find(|o| o.domain.is_empty() || o.domain == "ai.onnx")
.map_or(0, |o| o.version);
let input_names = graph
.input
.iter()
.filter(|input| !graph.initializer.iter().any(|t| t.name == input.name))
.map(|input| input.name.clone())
.collect();
let output_names = graph.output.iter().map(|o| o.name.clone()).collect();
Ok(GraphStats {
parameters,
nodes: graph.node.len() as u64,
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,
ir_version: 9,
opset: 17,
input_names: vec!["board".to_string()],
output_names: vec!["policy".to_string()],
}
);
}
#[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 {
output: vec!["c".to_string()],
op_type: "Constant".to_string(),
},
NodeProto {
output: vec!["y".to_string()],
op_type: "Add".to_string(),
},
],
initializer: vec![
TensorProto {
dims: vec![3, 4],
data_type: 1,
name: "w".to_string(),
},
TensorProto {
dims: vec![],
data_type: 1,
name: "scale".to_string(),
},
],
input: vec![
ValueInfoProto {
name: "x".to_string(),
},
ValueInfoProto {
name: "w".to_string(),
},
],
output: vec![ValueInfoProto {
name: "y".to_string(),
}],
}),
};
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 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,
opset_import: vec![],
graph: Some(GraphProto::default()),
};
let err = read_stats(&graph_only.encode_to_vec()).expect_err("no ir_version");
assert!(err.contains("IR version"), "{err}");
}
}