use crate::external_data::{ExternalDataInfo, ExternalDataLoader};
use crate::tensor::TensorDataLocation;
use crate::{
AttributeProto, AttributeValue, DataType, Error, Graph, GraphProto, NodeProto, Operation,
Tensor, TensorProto, TypeProto, tensor_shape_proto::dimension::Value, type_proto,
};
use std::collections::hash_map::{Entry, HashMap};
use std::sync::Arc;
pub(crate) fn tensor_from_proto(
tensor: TensorProto,
external_data_loader: Option<&Arc<ExternalDataLoader>>,
) -> Result<Tensor, Error> {
let data_type = DataType::from_onnx_type(tensor.data_type.unwrap_or(0));
let data = if !tensor.external_data.is_empty() {
let loader = external_data_loader.ok_or(Error::ExternalDataRequiresPath)?;
let external_info =
ExternalDataInfo::from_key_value_pairs(tensor.external_data, loader.clone())?;
TensorDataLocation::External(external_info)
} else if let Some(raw) = tensor.raw_data {
TensorDataLocation::Mmap(raw)
} else {
match data_type {
DataType::Undefined => TensorDataLocation::None,
DataType::String => TensorDataLocation::MmapStrings(tensor.string_data),
DataType::Float | DataType::Complex64 => TensorDataLocation::F32(tensor.float_data),
DataType::Double | DataType::Complex128 => TensorDataLocation::F64(tensor.double_data),
DataType::Int64 => TensorDataLocation::I64(tensor.int64_data),
DataType::Uint32 | DataType::Uint64 => TensorDataLocation::U64(tensor.uint64_data),
DataType::Int32
| DataType::Int16
| DataType::Int8
| DataType::Int4
| DataType::Int2
| DataType::Uint16
| DataType::Uint8
| DataType::Uint4
| DataType::Uint2
| DataType::Bool
| DataType::Float16
| DataType::Bfloat16
| DataType::Float8e4m3fn
| DataType::Float8e4m3fnuz
| DataType::Float8e5m2
| DataType::Float8e5m2fnuz
| DataType::Float8e8m0
| DataType::Float4e2m1 => TensorDataLocation::I32(tensor.int32_data),
}
};
Ok(Tensor::new(tensor.name, tensor.dims, data_type, data))
}
fn tensor_from_value_info(name: Option<String>, vi_type: Option<TypeProto>) -> Option<Tensor> {
let Some(type_proto::Value::TensorType(tensor_type)) = vi_type.and_then(|t| t.value) else {
return None;
};
let shape = tensor_type
.shape
.iter()
.flat_map(|s| &s.dim)
.map(|d| match d.value {
Some(Value::DimValue(v)) => v,
_ => -1,
})
.collect();
let data_type = DataType::from_onnx_type(tensor_type.elem_type.unwrap_or(0));
Some(Tensor::new(
name,
shape,
data_type,
TensorDataLocation::None,
))
}
pub(crate) fn graph_from_proto(
graph: GraphProto,
external_data_loader: Option<&Arc<ExternalDataLoader>>,
) -> Result<Graph, Error> {
let mut tensors = HashMap::with_capacity(
graph.initializer.len() + graph.value_info.len() + graph.input.len() + graph.output.len(),
);
for tensor in graph.initializer {
let onnx_tensor = tensor_from_proto(tensor, external_data_loader)?;
let name = onnx_tensor
.name()
.filter(|n| !n.is_empty())
.ok_or(Error::MissingField("initialiser tensor name"))?;
tensors.insert(name.to_string(), onnx_tensor);
}
let mut inputs = Vec::with_capacity(graph.input.len());
for input in graph.input {
let name = input
.name
.filter(|n| !n.is_empty())
.ok_or(Error::MissingField("graph input name"))?;
if let Entry::Vacant(entry) = tensors.entry(name) {
inputs.push(entry.key().clone());
if let Some(tensor) = tensor_from_value_info(Some(entry.key().clone()), input.r#type) {
entry.insert(tensor);
}
}
}
let mut outputs = Vec::with_capacity(graph.output.len());
for output in graph.output {
let name = output
.name
.filter(|n| !n.is_empty())
.ok_or(Error::MissingField("graph output name"))?;
outputs.push(name.clone());
if let Entry::Vacant(entry) = tensors.entry(name)
&& let Some(tensor) = tensor_from_value_info(Some(entry.key().clone()), output.r#type)
{
entry.insert(tensor);
}
}
for value_info in graph.value_info {
if let Some(name) = value_info.name.filter(|n| !n.is_empty())
&& let Entry::Vacant(entry) = tensors.entry(name)
&& let Some(tensor) =
tensor_from_value_info(Some(entry.key().clone()), value_info.r#type)
{
entry.insert(tensor);
}
}
let operations = graph
.node
.into_iter()
.map(|node| operation_from_node_proto(node, external_data_loader))
.collect::<Result<Vec<_>, Error>>()?;
Ok(Graph::new(graph.name, tensors, operations, inputs, outputs))
}
pub(crate) fn operation_from_node_proto(
node: NodeProto,
external_data_loader: Option<&Arc<ExternalDataLoader>>,
) -> Result<Operation, Error> {
let op_type = node
.op_type
.filter(|s| !s.is_empty())
.ok_or(Error::MissingField("node op_type"))?;
let attributes: HashMap<String, AttributeValue> = node
.attribute
.into_iter()
.map(|attr| parse_attribute_proto(attr, external_data_loader))
.collect::<Result<HashMap<_, _>, Error>>()?;
Ok(Operation::new(
node.name,
op_type,
node.input,
node.output,
attributes,
))
}
pub(crate) fn parse_attribute_proto(
attr: AttributeProto,
external_data_loader: Option<&Arc<ExternalDataLoader>>,
) -> Result<(String, AttributeValue), Error> {
let name = attr
.name
.filter(|n| !n.is_empty())
.ok_or(Error::MissingField("attribute name"))?;
let value = match attr.r#type.ok_or(Error::MissingField("attribute type"))? {
1 => Ok(AttributeValue::Float(attr.f.unwrap_or(0.0))),
2 => Ok(AttributeValue::Int(attr.i.unwrap_or(0))),
3 => Ok(AttributeValue::String(attr.s.unwrap_or_default())),
4 => {
let tensor = attr.t.ok_or(Error::MissingField("tensor attribute data"))?;
let onnx_tensor = tensor_from_proto(tensor, external_data_loader)?;
Ok(AttributeValue::Tensor(Box::new(onnx_tensor)))
}
5 => {
let graph = attr.g.ok_or(Error::MissingField("graph attribute data"))?;
let onnx_graph = graph_from_proto(graph, external_data_loader)?;
Ok(AttributeValue::Graph(Box::new(onnx_graph)))
}
6 => Ok(AttributeValue::Floats(attr.floats)),
7 => Ok(AttributeValue::Ints(attr.ints)),
8 => Ok(AttributeValue::Strings(attr.strings)),
9 => Ok(AttributeValue::Tensors(
attr.tensors
.into_iter()
.map(|tensor| tensor_from_proto(tensor, external_data_loader))
.collect::<Result<Box<[Tensor]>, Error>>()?,
)),
10 => Ok(AttributeValue::Graphs(
attr.graphs
.into_iter()
.map(|graph| graph_from_proto(graph, external_data_loader))
.collect::<Result<Box<[Graph]>, Error>>()?,
)),
n => Err(Error::UnsupportedAttributeType(n)),
}?;
Ok((name, value))
}