1use std::collections::HashMap;
28use std::fs;
29use std::path::Path;
30
31use shrew_core::backend::Backend;
32use shrew_core::dtype::DType;
33use shrew_core::error::Result;
34use shrew_core::tensor::Tensor;
35
36use shrew_nn::Module;
37use shrew_ir::graph::{Dim, DType as IrDType, IrGraph, IrType, OpKind};
38
39const ONNX_IR_VERSION: i64 = 9;
45const ONNX_OPSET_VERSION: i64 = 17;
47const ONNX_DOMAIN: &str = "";
49
50const ONNX_FLOAT: i32 = 1;
53const ONNX_DOUBLE: i32 = 11;
54const ONNX_FLOAT16: i32 = 10;
55const ONNX_BFLOAT16: i32 = 16;
56const ONNX_INT8: i32 = 3;
57const ONNX_UINT8: i32 = 2;
58const ONNX_INT32: i32 = 6;
59const ONNX_INT64: i32 = 7;
60const ONNX_UINT32: i32 = 12;
61
62struct PbEncoder {
71 buf: Vec<u8>,
72}
73
74impl PbEncoder {
75 fn new() -> Self {
76 Self { buf: Vec::new() }
77 }
78
79 fn into_bytes(self) -> Vec<u8> {
80 self.buf
81 }
82
83 fn write_varint(&mut self, mut val: u64) {
85 loop {
86 let byte = (val & 0x7F) as u8;
87 val >>= 7;
88 if val == 0 {
89 self.buf.push(byte);
90 break;
91 } else {
92 self.buf.push(byte | 0x80);
93 }
94 }
95 }
96
97 fn write_tag(&mut self, field: u32, wire_type: u32) {
99 self.write_varint(((field as u64) << 3) | wire_type as u64);
100 }
101
102 fn write_varint_field(&mut self, field: u32, val: u64) {
104 self.write_tag(field, 0);
105 self.write_varint(val);
106 }
107
108 fn write_sint64_field(&mut self, field: u32, val: i64) {
110 self.write_varint_field(field, val as u64);
111 }
112
113 fn write_bytes_field(&mut self, field: u32, data: &[u8]) {
115 self.write_tag(field, 2);
116 self.write_varint(data.len() as u64);
117 self.buf.extend_from_slice(data);
118 }
119
120 fn write_string_field(&mut self, field: u32, val: &str) {
122 self.write_bytes_field(field, val.as_bytes());
123 }
124
125 fn write_message_field(&mut self, field: u32, encoder: &PbEncoder) {
127 self.write_bytes_field(field, &encoder.buf);
128 }
129
130 #[allow(dead_code)]
132 fn write_float_data(&mut self, field: u32, data: &[f32]) {
133 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_le_bytes()).collect();
134 self.write_bytes_field(field, &bytes);
135 }
136
137 #[allow(dead_code)]
139 fn write_double_data(&mut self, field: u32, data: &[f64]) {
140 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_le_bytes()).collect();
141 self.write_bytes_field(field, &bytes);
142 }
143}
144
145struct PbDecoder<'a> {
151 data: &'a [u8],
152 pos: usize,
153}
154
155impl<'a> PbDecoder<'a> {
156 fn new(data: &'a [u8]) -> Self {
157 Self { data, pos: 0 }
158 }
159
160 fn remaining(&self) -> usize {
161 self.data.len() - self.pos
162 }
163
164 fn read_varint(&mut self) -> Result<u64> {
165 let mut result: u64 = 0;
166 let mut shift = 0;
167 loop {
168 if self.pos >= self.data.len() {
169 return Err(shrew_core::Error::msg("protobuf: unexpected end of data"));
170 }
171 let byte = self.data[self.pos];
172 self.pos += 1;
173 result |= ((byte & 0x7F) as u64) << shift;
174 if byte & 0x80 == 0 {
175 break;
176 }
177 shift += 7;
178 if shift > 63 {
179 return Err(shrew_core::Error::msg("protobuf: varint too long"));
180 }
181 }
182 Ok(result)
183 }
184
185 fn read_tag(&mut self) -> Result<(u32, u32)> {
186 let val = self.read_varint()?;
187 let field = (val >> 3) as u32;
188 let wire_type = (val & 0x7) as u32;
189 Ok((field, wire_type))
190 }
191
192 fn read_bytes(&mut self) -> Result<&'a [u8]> {
193 let len = self.read_varint()? as usize;
194 if self.pos + len > self.data.len() {
195 return Err(shrew_core::Error::msg("protobuf: bytes field exceeds data"));
196 }
197 let result = &self.data[self.pos..self.pos + len];
198 self.pos += len;
199 Ok(result)
200 }
201
202 fn read_string(&mut self) -> Result<String> {
203 let bytes = self.read_bytes()?;
204 String::from_utf8(bytes.to_vec())
205 .map_err(|_| shrew_core::Error::msg("protobuf: invalid UTF-8 string"))
206 }
207
208 fn skip_field(&mut self, wire_type: u32) -> Result<()> {
209 match wire_type {
210 0 => {
211 self.read_varint()?;
212 }
213 1 => {
214 self.pos += 8;
215 } 2 => {
217 self.read_bytes()?;
218 }
219 5 => {
220 self.pos += 4;
221 } _ => {
223 return Err(shrew_core::Error::msg(format!(
224 "protobuf: unsupported wire type {wire_type}"
225 )))
226 }
227 }
228 Ok(())
229 }
230}
231
232#[derive(Debug, Clone)]
238pub struct OnnxTensor {
239 pub name: String,
241 pub data_type: i32,
243 pub dims: Vec<i64>,
245 pub float_data: Vec<f32>,
247 pub double_data: Vec<f64>,
249 pub raw_data: Vec<u8>,
251}
252
253impl OnnxTensor {
254 fn new(name: &str) -> Self {
255 Self {
256 name: name.to_string(),
257 data_type: ONNX_FLOAT,
258 dims: Vec::new(),
259 float_data: Vec::new(),
260 double_data: Vec::new(),
261 raw_data: Vec::new(),
262 }
263 }
264
265 fn encode(&self) -> Vec<u8> {
267 let mut enc = PbEncoder::new();
268 for &d in &self.dims {
270 enc.write_sint64_field(1, d);
271 }
272 enc.write_varint_field(2, self.data_type as u64);
274 if !self.name.is_empty() {
276 enc.write_string_field(8, &self.name);
277 }
278 if !self.float_data.is_empty() {
280 let bytes: Vec<u8> = self
282 .float_data
283 .iter()
284 .flat_map(|v| v.to_le_bytes())
285 .collect();
286 enc.write_bytes_field(13, &bytes);
287 } else if !self.double_data.is_empty() {
288 let bytes: Vec<u8> = self
289 .double_data
290 .iter()
291 .flat_map(|v| v.to_le_bytes())
292 .collect();
293 enc.write_bytes_field(13, &bytes);
294 } else if !self.raw_data.is_empty() {
295 enc.write_bytes_field(13, &self.raw_data);
296 }
297 enc.into_bytes()
298 }
299
300 fn decode(data: &[u8]) -> Result<Self> {
302 let mut dec = PbDecoder::new(data);
303 let mut tensor = OnnxTensor::new("");
304 while dec.remaining() > 0 {
305 let (field, wire_type) = dec.read_tag()?;
306 match (field, wire_type) {
307 (1, 0) => {
308 let v = dec.read_varint()? as i64;
310 tensor.dims.push(v);
311 }
312 (1, 2) => {
313 let bytes = dec.read_bytes()?;
315 let mut sub = PbDecoder::new(bytes);
316 while sub.remaining() > 0 {
317 tensor.dims.push(sub.read_varint()? as i64);
318 }
319 }
320 (2, 0) => {
321 tensor.data_type = dec.read_varint()? as i32;
323 }
324 (8, 2) => {
325 tensor.name = dec.read_string()?;
327 }
328 (13, 2) => {
329 tensor.raw_data = dec.read_bytes()?.to_vec();
331 }
332 (4, 2) => {
333 let bytes = dec.read_bytes()?;
335 for chunk in bytes.chunks_exact(4) {
336 let val = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
337 tensor.float_data.push(val);
338 }
339 }
340 (4, 5) => {
341 let bytes = &dec.data[dec.pos..dec.pos + 4];
343 dec.pos += 4;
344 let val = f32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
345 tensor.float_data.push(val);
346 }
347 (5, 2) => {
348 let bytes = dec.read_bytes()?;
350 for chunk in bytes.chunks_exact(8) {
351 let val = f64::from_le_bytes([
352 chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6],
353 chunk[7],
354 ]);
355 tensor.double_data.push(val);
356 }
357 }
358 _ => {
359 dec.skip_field(wire_type)?;
360 }
361 }
362 }
363 Ok(tensor)
364 }
365
366 fn to_f64_vec(&self) -> Vec<f64> {
368 if !self.double_data.is_empty() {
369 return self.double_data.clone();
370 }
371 if !self.float_data.is_empty() {
372 return self.float_data.iter().map(|&v| v as f64).collect();
373 }
374 if !self.raw_data.is_empty() {
375 match self.data_type {
376 ONNX_FLOAT => self
377 .raw_data
378 .chunks_exact(4)
379 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]) as f64)
380 .collect(),
381 ONNX_DOUBLE => self
382 .raw_data
383 .chunks_exact(8)
384 .map(|c| f64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
385 .collect(),
386 ONNX_FLOAT16 => self
387 .raw_data
388 .chunks_exact(2)
389 .map(|c| {
390 let bits = u16::from_le_bytes([c[0], c[1]]);
391 half::f16::from_bits(bits).to_f64()
392 })
393 .collect(),
394 _ => Vec::new(),
395 }
396 } else {
397 Vec::new()
398 }
399 }
400}
401
402fn dtype_to_onnx(dtype: DType) -> i32 {
404 match dtype {
405 DType::F32 => ONNX_FLOAT,
406 DType::F64 => ONNX_DOUBLE,
407 DType::F16 => ONNX_FLOAT16,
408 DType::BF16 => ONNX_BFLOAT16,
409 DType::U8 => ONNX_UINT8,
410 DType::U32 => ONNX_UINT32,
411 DType::I64 => ONNX_INT64,
412 }
413}
414
415fn onnx_to_dtype(onnx_type: i32) -> Result<DType> {
417 match onnx_type {
418 ONNX_FLOAT => Ok(DType::F32),
419 ONNX_DOUBLE => Ok(DType::F64),
420 ONNX_FLOAT16 => Ok(DType::F16),
421 ONNX_BFLOAT16 => Ok(DType::BF16),
422 ONNX_UINT8 => Ok(DType::U8),
423 ONNX_UINT32 => Ok(DType::U32),
424 ONNX_INT64 => Ok(DType::I64),
425 ONNX_INT8 => Ok(DType::U8), ONNX_INT32 => Ok(DType::I64), _ => Err(shrew_core::Error::msg(format!(
428 "unsupported ONNX data type: {onnx_type}"
429 ))),
430 }
431}
432
433#[derive(Debug, Clone)]
439pub struct OnnxNode {
440 pub inputs: Vec<String>,
442 pub outputs: Vec<String>,
444 pub op_type: String,
446 pub name: String,
448 pub attributes: HashMap<String, OnnxAttribute>,
450}
451
452#[derive(Debug, Clone)]
454pub enum OnnxAttribute {
455 Int(i64),
456 Float(f32),
457 String(String),
458 Ints(Vec<i64>),
459 Floats(Vec<f32>),
460}
461
462impl OnnxNode {
463 fn encode(&self) -> Vec<u8> {
464 let mut enc = PbEncoder::new();
465 for input in &self.inputs {
467 enc.write_string_field(1, input);
468 }
469 for output in &self.outputs {
471 enc.write_string_field(2, output);
472 }
473 if !self.name.is_empty() {
475 enc.write_string_field(3, &self.name);
476 }
477 enc.write_string_field(4, &self.op_type);
479 for (key, val) in &self.attributes {
481 let attr = encode_attribute(key, val);
482 enc.write_message_field(5, &attr);
483 }
484 enc.into_bytes()
485 }
486}
487
488fn encode_attribute(name: &str, val: &OnnxAttribute) -> PbEncoder {
489 let mut enc = PbEncoder::new();
490 enc.write_string_field(1, name); match val {
492 OnnxAttribute::Int(i) => {
493 enc.write_varint_field(2, 2); enc.write_sint64_field(3, *i); }
496 OnnxAttribute::Float(f) => {
497 enc.write_varint_field(2, 1); enc.write_tag(4, 5);
500 enc.buf.extend_from_slice(&f.to_le_bytes());
501 }
502 OnnxAttribute::String(s) => {
503 enc.write_varint_field(2, 3); enc.write_bytes_field(5, s.as_bytes()); }
506 OnnxAttribute::Ints(ints) => {
507 enc.write_varint_field(2, 7); for &i in ints {
509 enc.write_sint64_field(8, i); }
511 }
512 OnnxAttribute::Floats(floats) => {
513 enc.write_varint_field(2, 6); for &f in floats {
515 enc.write_tag(7, 5); enc.buf.extend_from_slice(&f.to_le_bytes());
517 }
518 }
519 }
520 enc
521}
522
523#[derive(Debug, Clone)]
529pub struct OnnxModel {
530 pub producer_name: String,
532 pub producer_version: String,
534 pub graph_name: String,
536 pub nodes: Vec<OnnxNode>,
538 pub initializers: Vec<OnnxTensor>,
540 pub inputs: Vec<(String, Vec<i64>, i32)>,
542 pub outputs: Vec<(String, Vec<i64>, i32)>,
544}
545
546impl OnnxModel {
547 pub fn new(graph_name: &str) -> Self {
549 Self {
550 producer_name: "Shrew".to_string(),
551 producer_version: "0.1.0".to_string(),
552 graph_name: graph_name.to_string(),
553 nodes: Vec::new(),
554 initializers: Vec::new(),
555 inputs: Vec::new(),
556 outputs: Vec::new(),
557 }
558 }
559
560 pub fn to_bytes(&self) -> Vec<u8> {
562 let mut graph = PbEncoder::new();
564
565 for node in &self.nodes {
567 let node_bytes = node.encode();
568 graph.write_bytes_field(1, &node_bytes);
569 }
570
571 graph.write_string_field(2, &self.graph_name);
573
574 for init in &self.initializers {
576 let tensor_bytes = init.encode();
577 graph.write_bytes_field(5, &tensor_bytes);
578 }
579
580 for (name, dims, dtype) in &self.inputs {
582 let vi = encode_value_info(name, dims, *dtype);
583 graph.write_message_field(11, &vi);
584 }
585
586 for (name, dims, dtype) in &self.outputs {
588 let vi = encode_value_info(name, dims, *dtype);
589 graph.write_message_field(12, &vi);
590 }
591
592 let mut model = PbEncoder::new();
594 model.write_varint_field(1, ONNX_IR_VERSION as u64);
596 model.write_string_field(2, &self.producer_name);
598 model.write_string_field(3, &self.producer_version);
600 model.write_message_field(7, &graph);
602 let mut opset = PbEncoder::new();
604 opset.write_string_field(1, ONNX_DOMAIN); opset.write_varint_field(2, ONNX_OPSET_VERSION as u64); model.write_message_field(8, &opset);
607
608 model.into_bytes()
609 }
610
611 pub fn save<P: AsRef<Path>>(&self, path: P) -> Result<()> {
613 let bytes = self.to_bytes();
614 fs::write(path.as_ref(), &bytes)
615 .map_err(|e| shrew_core::Error::msg(format!("failed to write ONNX file: {e}")))
616 }
617}
618
619fn encode_value_info(name: &str, dims: &[i64], data_type: i32) -> PbEncoder {
621 let mut vi = PbEncoder::new();
622 vi.write_string_field(1, name); let mut type_proto = PbEncoder::new();
626 let mut tensor_type = PbEncoder::new();
628 tensor_type.write_varint_field(1, data_type as u64); let mut shape = PbEncoder::new();
631 for &d in dims {
632 let mut dim = PbEncoder::new();
633 if d >= 0 {
634 dim.write_sint64_field(1, d); } else {
636 dim.write_string_field(2, "dynamic"); }
638 shape.write_message_field(1, &dim);
639 }
640 tensor_type.write_message_field(2, &shape);
641 type_proto.write_message_field(1, &tensor_type);
642 vi.write_message_field(2, &type_proto);
643
644 vi
645}
646
647pub fn export_weights<P, B, M>(
669 path: P,
670 module: &M,
671 model_name: &str,
672 input_shape: &[i64],
673) -> Result<()>
674where
675 P: AsRef<Path>,
676 B: Backend,
677 M: Module<B>,
678{
679 let named = module.named_parameters();
680
681 let mut model = OnnxModel::new(model_name);
682
683 model
685 .inputs
686 .push(("input".to_string(), input_shape.to_vec(), ONNX_FLOAT));
687
688 for (name, tensor) in &named {
690 let data = tensor.to_f64_vec()?;
691 let dims: Vec<i64> = tensor.dims().iter().map(|&d| d as i64).collect();
692
693 let mut onnx_tensor = OnnxTensor::new(name);
694 onnx_tensor.data_type = dtype_to_onnx(tensor.dtype());
695 onnx_tensor.dims = dims;
696
697 match tensor.dtype() {
698 DType::F32 => {
699 onnx_tensor.float_data = data.iter().map(|&v| v as f32).collect();
700 }
701 DType::F64 => {
702 onnx_tensor.double_data = data;
703 }
704 _ => {
705 onnx_tensor.data_type = ONNX_FLOAT;
707 onnx_tensor.float_data = data.iter().map(|&v| v as f32).collect();
708 }
709 }
710
711 model.initializers.push(onnx_tensor);
712 }
713
714 model.outputs.push((
717 "output".to_string(),
718 vec![-1], ONNX_FLOAT,
720 ));
721
722 model.save(path)
723}
724
725pub fn export_tensors<P, B>(
729 path: P,
730 tensors: &[(String, Tensor<B>)],
731 model_name: &str,
732) -> Result<()>
733where
734 P: AsRef<Path>,
735 B: Backend,
736{
737 let mut model = OnnxModel::new(model_name);
738
739 for (name, tensor) in tensors {
740 let data = tensor.to_f64_vec()?;
741 let dims: Vec<i64> = tensor.dims().iter().map(|&d| d as i64).collect();
742
743 let mut onnx_tensor = OnnxTensor::new(name);
744 onnx_tensor.data_type = dtype_to_onnx(tensor.dtype());
745 onnx_tensor.dims = dims;
746
747 match tensor.dtype() {
748 DType::F32 => {
749 onnx_tensor.float_data = data.iter().map(|&v| v as f32).collect();
750 }
751 DType::F64 => {
752 onnx_tensor.double_data = data;
753 }
754 _ => {
755 onnx_tensor.data_type = ONNX_FLOAT;
756 onnx_tensor.float_data = data.iter().map(|&v| v as f32).collect();
757 }
758 }
759
760 model.initializers.push(onnx_tensor);
761 }
762
763 model.save(path)
764}
765
766pub fn ir_graph_to_onnx(graph: &IrGraph) -> Result<OnnxModel> {
771 let mut model = OnnxModel::new(&graph.name);
772
773 for &input_id in &graph.inputs {
775 let node = &graph.nodes[input_id.0];
776 let (dims, dtype) = match &node.output_type {
777 IrType::Tensor { shape, dtype } => {
778 let d: Vec<i64> = shape
779 .iter()
780 .map(|dim| match dim {
781 Dim::Fixed(n) => *n,
782 _ => -1,
783 })
784 .collect();
785 let onnx_dt = match dtype {
786 IrDType::F32 => ONNX_FLOAT,
787 IrDType::F64 => ONNX_DOUBLE,
788 IrDType::F16 => ONNX_FLOAT16,
789 IrDType::Bf16 => ONNX_BFLOAT16,
790 IrDType::I8 => ONNX_INT8,
791 IrDType::U8 => ONNX_UINT8,
792 IrDType::I32 => ONNX_INT32,
793 IrDType::I64 => ONNX_INT64,
794 _ => ONNX_FLOAT,
795 };
796 (d, onnx_dt)
797 }
798 _ => (vec![-1], ONNX_FLOAT),
799 };
800 model.inputs.push((node.name.clone(), dims, dtype));
801 }
802
803 for out in &graph.outputs {
805 let node = &graph.nodes[out.node_id.0];
806 let (dims, dtype) = match &node.output_type {
807 IrType::Tensor { shape, dtype } => {
808 let d: Vec<i64> = shape
809 .iter()
810 .map(|dim| match dim {
811 Dim::Fixed(n) => *n,
812 _ => -1,
813 })
814 .collect();
815 let onnx_dt = match dtype {
816 IrDType::F32 => ONNX_FLOAT,
817 IrDType::F64 => ONNX_DOUBLE,
818 IrDType::F16 => ONNX_FLOAT16,
819 IrDType::Bf16 => ONNX_BFLOAT16,
820 IrDType::I8 => ONNX_INT8,
821 IrDType::U8 => ONNX_UINT8,
822 IrDType::I32 => ONNX_INT32,
823 IrDType::I64 => ONNX_INT64,
824 _ => ONNX_FLOAT,
825 };
826 (d, onnx_dt)
827 }
828 _ => (vec![-1], ONNX_FLOAT),
829 };
830 model.outputs.push((out.name.clone(), dims, dtype));
831 }
832
833 for node in &graph.nodes {
835 if graph.inputs.contains(&node.id) {
836 continue;
837 }
838
839 let input_names: Vec<String> = node
840 .inputs
841 .iter()
842 .map(|id| graph.nodes[id.0].name.clone())
843 .collect();
844 let output_names = vec![node.name.clone()];
845 let mut attributes = HashMap::new();
846
847 let op_type = match &node.op {
848 OpKind::Add => "Add",
849 OpKind::Sub => "Sub",
850 OpKind::Mul => "Mul",
851 OpKind::Div => "Div",
852 OpKind::Pow => "Pow",
853 OpKind::MatMul => "MatMul",
854 OpKind::Relu => "Relu",
855 OpKind::Gelu => "Gelu",
856 OpKind::Silu => "Silu",
857 OpKind::Sigmoid => "Sigmoid",
858 OpKind::Tanh => "Tanh",
859 OpKind::Exp => "Exp",
860 OpKind::Log => "Log",
861 OpKind::Sqrt => "Sqrt",
862 OpKind::Neg => "Neg",
863 OpKind::Transpose => "Transpose",
864 OpKind::Reshape { .. } => "Reshape",
865 OpKind::Identity => "Identity",
866 OpKind::Equal => "Equal",
867 OpKind::NotEqual => "NotEqual",
868 OpKind::Less => "Less",
869 OpKind::Greater => "Greater",
870 OpKind::LessEqual => "LessOrEqual",
871 OpKind::GreaterEqual => "GreaterOrEqual",
872 OpKind::And => "And",
873 OpKind::Or => "Or",
874 OpKind::Not => "Not",
875 OpKind::Softmax { dim } => {
876 attributes.insert("axis".to_string(), OnnxAttribute::Int(*dim));
877 "Softmax"
878 }
879 OpKind::Sum { dims, keepdim } => {
880 attributes.insert("axes".to_string(), OnnxAttribute::Ints(dims.clone()));
881 attributes.insert("keepdims".to_string(), OnnxAttribute::Int(*keepdim as i64));
882 "ReduceSum"
883 }
884 OpKind::Mean { dims, keepdim } => {
885 attributes.insert("axes".to_string(), OnnxAttribute::Ints(dims.clone()));
886 attributes.insert("keepdims".to_string(), OnnxAttribute::Int(*keepdim as i64));
887 "ReduceMean"
888 }
889 OpKind::Max { dim, keepdim } => {
890 attributes.insert("axes".to_string(), OnnxAttribute::Ints(vec![*dim]));
891 attributes.insert("keepdims".to_string(), OnnxAttribute::Int(*keepdim as i64));
892 "ReduceMax"
893 }
894 OpKind::Min { dim, keepdim } => {
895 attributes.insert("axes".to_string(), OnnxAttribute::Ints(vec![*dim]));
896 attributes.insert("keepdims".to_string(), OnnxAttribute::Int(*keepdim as i64));
897 "ReduceMin"
898 }
899 OpKind::LayerNorm { eps } => {
900 attributes.insert("epsilon".to_string(), OnnxAttribute::Float(*eps as f32));
901 "LayerNormalization"
902 }
903 OpKind::BatchNorm { eps } => {
904 attributes.insert("epsilon".to_string(), OnnxAttribute::Float(*eps as f32));
905 "BatchNormalization"
906 }
907 OpKind::Linear { .. } => "Gemm",
908 OpKind::Dropout { p } => {
909 attributes.insert("ratio".to_string(), OnnxAttribute::Float(*p as f32));
910 "Dropout"
911 }
912 OpKind::Concat { dim } => {
913 attributes.insert("axis".to_string(), OnnxAttribute::Int(*dim));
914 "Concat"
915 }
916 OpKind::Permute { dims } => {
917 attributes.insert("perm".to_string(), OnnxAttribute::Ints(dims.clone()));
918 "Transpose"
919 }
920 _ => &node.name,
921 };
922
923 model.nodes.push(OnnxNode {
924 inputs: input_names,
925 outputs: output_names,
926 op_type: op_type.to_string(),
927 name: node.name.clone(),
928 attributes,
929 });
930 }
931
932 Ok(model)
933}
934
935pub fn export_ir_graph<P: AsRef<Path>>(path: P, graph: &IrGraph) -> Result<()> {
937 let model = ir_graph_to_onnx(graph)?;
938 model.save(path)
939}
940
941pub fn load_onnx_weights<B: Backend>(
957 path: impl AsRef<Path>,
958 device: &B::Device,
959) -> Result<HashMap<String, Tensor<B>>> {
960 let bytes = fs::read(path.as_ref())
961 .map_err(|e| shrew_core::Error::msg(format!("failed to read ONNX file: {e}")))?;
962
963 load_onnx_weights_from_bytes::<B>(&bytes, device)
964}
965
966pub fn load_onnx_weights_from_bytes<B: Backend>(
968 data: &[u8],
969 device: &B::Device,
970) -> Result<HashMap<String, Tensor<B>>> {
971 let mut dec = PbDecoder::new(data);
972 let mut result = HashMap::new();
973
974 while dec.remaining() > 0 {
976 let (field, wire_type) = dec.read_tag()?;
977 match (field, wire_type) {
978 (7, 2) => {
979 let graph_bytes = dec.read_bytes()?;
981 let tensors = parse_graph_initializers::<B>(graph_bytes, device)?;
982 result.extend(tensors);
983 }
984 _ => {
985 dec.skip_field(wire_type)?;
986 }
987 }
988 }
989
990 Ok(result)
991}
992
993fn parse_graph_initializers<B: Backend>(
995 data: &[u8],
996 device: &B::Device,
997) -> Result<HashMap<String, Tensor<B>>> {
998 let mut dec = PbDecoder::new(data);
999 let mut result = HashMap::new();
1000
1001 while dec.remaining() > 0 {
1002 let (field, wire_type) = dec.read_tag()?;
1003 match (field, wire_type) {
1004 (5, 2) => {
1005 let tensor_bytes = dec.read_bytes()?;
1007 let onnx_tensor = OnnxTensor::decode(tensor_bytes)?;
1008
1009 if !onnx_tensor.name.is_empty() {
1010 let dtype = onnx_to_dtype(onnx_tensor.data_type)?;
1011 let shape: Vec<usize> = onnx_tensor.dims.iter().map(|&d| d as usize).collect();
1012 let f64_data = onnx_tensor.to_f64_vec();
1013
1014 if !f64_data.is_empty() {
1015 let tensor = Tensor::<B>::from_f64_slice(&f64_data, shape, dtype, device)?;
1016 result.insert(onnx_tensor.name.clone(), tensor);
1017 }
1018 }
1019 }
1020 _ => {
1021 dec.skip_field(wire_type)?;
1022 }
1023 }
1024 }
1025
1026 Ok(result)
1027}
1028
1029#[derive(Debug, Clone)]
1035pub struct OnnxGraph {
1036 pub nodes: Vec<OnnxNode>,
1038 pub initializer_protos: Vec<OnnxTensor>,
1040 pub input_names: Vec<String>,
1042 pub output_names: Vec<String>,
1044 pub name: String,
1046}
1047
1048fn decode_node(data: &[u8]) -> Result<OnnxNode> {
1050 let mut dec = PbDecoder::new(data);
1051 let mut node = OnnxNode {
1052 inputs: Vec::new(),
1053 outputs: Vec::new(),
1054 op_type: String::new(),
1055 name: String::new(),
1056 attributes: HashMap::new(),
1057 };
1058 while dec.remaining() > 0 {
1059 let (field, wire_type) = dec.read_tag()?;
1060 match (field, wire_type) {
1061 (1, 2) => node.inputs.push(dec.read_string()?),
1062 (2, 2) => node.outputs.push(dec.read_string()?),
1063 (3, 2) => node.name = dec.read_string()?,
1064 (4, 2) => node.op_type = dec.read_string()?,
1065 (5, 2) => {
1066 let attr_bytes = dec.read_bytes()?;
1067 let (key, val) = decode_attribute(attr_bytes)?;
1068 node.attributes.insert(key, val);
1069 }
1070 _ => dec.skip_field(wire_type)?,
1071 }
1072 }
1073 Ok(node)
1074}
1075
1076fn decode_attribute(data: &[u8]) -> Result<(String, OnnxAttribute)> {
1078 let mut dec = PbDecoder::new(data);
1079 let mut name = String::new();
1080 let mut attr_type: u64 = 0;
1081 let mut int_val: i64 = 0;
1082 let mut float_val: f32 = 0.0;
1083 let mut string_val = Vec::new();
1084 let mut ints_val: Vec<i64> = Vec::new();
1085 let mut floats_val: Vec<f32> = Vec::new();
1086 while dec.remaining() > 0 {
1087 let (field, wire_type) = dec.read_tag()?;
1088 match (field, wire_type) {
1089 (1, 2) => name = dec.read_string()?, (2, 0) => attr_type = dec.read_varint()?, (3, 0) => int_val = dec.read_varint()? as i64, (4, 5) => {
1093 if dec.pos + 4 > dec.data.len() {
1095 return Err(shrew_core::Error::msg("attribute: unexpected end"));
1096 }
1097 let b = &dec.data[dec.pos..dec.pos + 4];
1098 float_val = f32::from_le_bytes([b[0], b[1], b[2], b[3]]);
1099 dec.pos += 4;
1100 }
1101 (5, 2) => string_val = dec.read_bytes()?.to_vec(), (7, 5) => {
1103 if dec.pos + 4 > dec.data.len() {
1105 return Err(shrew_core::Error::msg("attribute: unexpected end"));
1106 }
1107 let b = &dec.data[dec.pos..dec.pos + 4];
1108 floats_val.push(f32::from_le_bytes([b[0], b[1], b[2], b[3]]));
1109 dec.pos += 4;
1110 }
1111 (7, 2) => {
1112 let bytes = dec.read_bytes()?;
1114 for c in bytes.chunks_exact(4) {
1115 floats_val.push(f32::from_le_bytes([c[0], c[1], c[2], c[3]]));
1116 }
1117 }
1118 (8, 0) => ints_val.push(dec.read_varint()? as i64), (8, 2) => {
1120 let bytes = dec.read_bytes()?;
1122 let mut sub = PbDecoder::new(bytes);
1123 while sub.remaining() > 0 {
1124 ints_val.push(sub.read_varint()? as i64);
1125 }
1126 }
1127 _ => dec.skip_field(wire_type)?,
1128 }
1129 }
1130 let val = match attr_type {
1131 1 => OnnxAttribute::Float(float_val),
1132 2 => OnnxAttribute::Int(int_val),
1133 3 => OnnxAttribute::String(String::from_utf8(string_val).unwrap_or_default()),
1134 6 => OnnxAttribute::Floats(floats_val),
1135 7 => OnnxAttribute::Ints(ints_val),
1136 _ => OnnxAttribute::Int(int_val), };
1138 Ok((name, val))
1139}
1140
1141fn parse_graph_proto(data: &[u8]) -> Result<OnnxGraph> {
1143 let mut dec = PbDecoder::new(data);
1144 let mut graph = OnnxGraph {
1145 nodes: Vec::new(),
1146 initializer_protos: Vec::new(),
1147 input_names: Vec::new(),
1148 output_names: Vec::new(),
1149 name: String::new(),
1150 };
1151 while dec.remaining() > 0 {
1152 let (field, wire_type) = dec.read_tag()?;
1153 match (field, wire_type) {
1154 (1, 2) => {
1155 let node_bytes = dec.read_bytes()?;
1156 graph.nodes.push(decode_node(node_bytes)?);
1157 }
1158 (2, 2) => graph.name = dec.read_string()?,
1159 (5, 2) => {
1160 let tensor_bytes = dec.read_bytes()?;
1161 graph
1162 .initializer_protos
1163 .push(OnnxTensor::decode(tensor_bytes)?);
1164 }
1165 (11, 2) => {
1166 let vi_bytes = dec.read_bytes()?;
1168 let name = extract_value_info_name(vi_bytes)?;
1169 graph.input_names.push(name);
1170 }
1171 (12, 2) => {
1172 let vi_bytes = dec.read_bytes()?;
1174 let name = extract_value_info_name(vi_bytes)?;
1175 graph.output_names.push(name);
1176 }
1177 _ => dec.skip_field(wire_type)?,
1178 }
1179 }
1180 Ok(graph)
1181}
1182
1183fn extract_value_info_name(data: &[u8]) -> Result<String> {
1185 let mut dec = PbDecoder::new(data);
1186 while dec.remaining() > 0 {
1187 let (field, wire_type) = dec.read_tag()?;
1188 if field == 1 && wire_type == 2 {
1189 return dec.read_string();
1190 }
1191 dec.skip_field(wire_type)?;
1192 }
1193 Ok(String::new())
1194}
1195
1196pub fn load_onnx_graph(path: impl AsRef<Path>) -> Result<OnnxGraph> {
1198 let bytes = fs::read(path.as_ref())
1199 .map_err(|e| shrew_core::Error::msg(format!("failed to read ONNX file: {e}")))?;
1200 load_onnx_graph_from_bytes(&bytes)
1201}
1202
1203pub fn load_onnx_graph_from_bytes(data: &[u8]) -> Result<OnnxGraph> {
1205 let mut dec = PbDecoder::new(data);
1206 while dec.remaining() > 0 {
1207 let (field, wire_type) = dec.read_tag()?;
1208 if field == 7 && wire_type == 2 {
1209 let graph_bytes = dec.read_bytes()?;
1210 return parse_graph_proto(graph_bytes);
1211 }
1212 dec.skip_field(wire_type)?;
1213 }
1214 Err(shrew_core::Error::msg("ONNX file contains no graph"))
1215}
1216
1217pub fn run_onnx_graph<B: Backend>(
1239 graph: &OnnxGraph,
1240 inputs: &HashMap<String, Tensor<B>>,
1241 device: &B::Device,
1242) -> Result<HashMap<String, Tensor<B>>> {
1243 let mut env: HashMap<String, Tensor<B>> = HashMap::new();
1244
1245 for init in &graph.initializer_protos {
1247 if init.name.is_empty() {
1248 continue;
1249 }
1250 let dtype = onnx_to_dtype(init.data_type)?;
1251 let shape: Vec<usize> = init.dims.iter().map(|&d| d as usize).collect();
1252 let f64_data = init.to_f64_vec();
1253 if !f64_data.is_empty() {
1254 let tensor = Tensor::<B>::from_f64_slice(&f64_data, shape, dtype, device)?;
1255 env.insert(init.name.clone(), tensor);
1256 }
1257 }
1258
1259 for (name, tensor) in inputs {
1261 env.insert(name.clone(), tensor.clone());
1262 }
1263
1264 for node in &graph.nodes {
1266 execute_node(node, &mut env, device)?;
1267 }
1268
1269 let mut outputs = HashMap::new();
1271 for name in &graph.output_names {
1272 if let Some(t) = env.get(name) {
1273 outputs.insert(name.clone(), t.clone());
1274 }
1275 }
1276 Ok(outputs)
1277}
1278
1279fn get_tensor<'a, B: Backend>(
1281 env: &'a HashMap<String, Tensor<B>>,
1282 name: &str,
1283) -> Result<&'a Tensor<B>> {
1284 env.get(name)
1285 .ok_or_else(|| shrew_core::Error::msg(format!("ONNX runtime: tensor '{name}' not found")))
1286}
1287
1288fn attr_i(node: &OnnxNode, key: &str, default: i64) -> i64 {
1290 match node.attributes.get(key) {
1291 Some(OnnxAttribute::Int(v)) => *v,
1292 _ => default,
1293 }
1294}
1295
1296fn attr_ints(node: &OnnxNode, key: &str) -> Vec<i64> {
1298 match node.attributes.get(key) {
1299 Some(OnnxAttribute::Ints(v)) => v.clone(),
1300 _ => Vec::new(),
1301 }
1302}
1303
1304fn attr_f(node: &OnnxNode, key: &str, default: f32) -> f32 {
1306 match node.attributes.get(key) {
1307 Some(OnnxAttribute::Float(v)) => *v,
1308 _ => default,
1309 }
1310}
1311
1312fn execute_node<B: Backend>(
1314 node: &OnnxNode,
1315 env: &mut HashMap<String, Tensor<B>>,
1316 device: &B::Device,
1317) -> Result<()> {
1318 match node.op_type.as_str() {
1319 "Add" => {
1321 let a = get_tensor(env, &node.inputs[0])?;
1322 let b = get_tensor(env, &node.inputs[1])?;
1323 let out = a.add(b)?;
1324 env.insert(node.outputs[0].clone(), out);
1325 }
1326 "Sub" => {
1327 let a = get_tensor(env, &node.inputs[0])?;
1328 let b = get_tensor(env, &node.inputs[1])?;
1329 let out = a.sub(b)?;
1330 env.insert(node.outputs[0].clone(), out);
1331 }
1332 "Mul" => {
1333 let a = get_tensor(env, &node.inputs[0])?;
1334 let b = get_tensor(env, &node.inputs[1])?;
1335 let out = a.mul(b)?;
1336 env.insert(node.outputs[0].clone(), out);
1337 }
1338 "Div" => {
1339 let a = get_tensor(env, &node.inputs[0])?;
1340 let b = get_tensor(env, &node.inputs[1])?;
1341 let out = a.div(b)?;
1342 env.insert(node.outputs[0].clone(), out);
1343 }
1344 "Pow" => {
1345 let a = get_tensor(env, &node.inputs[0])?;
1346 let b = get_tensor(env, &node.inputs[1])?;
1348 let exp_val = b.to_f64_vec()?;
1349 if exp_val.len() == 1 {
1350 let out = a.powf(exp_val[0])?;
1351 env.insert(node.outputs[0].clone(), out);
1352 } else {
1353 return Err(shrew_core::Error::msg(
1354 "ONNX Pow: only scalar exponent supported",
1355 ));
1356 }
1357 }
1358
1359 "MatMul" => {
1361 let a = get_tensor(env, &node.inputs[0])?;
1362 let b = get_tensor(env, &node.inputs[1])?;
1363 let out = a.matmul(b)?;
1364 env.insert(node.outputs[0].clone(), out);
1365 }
1366 "Gemm" => {
1367 let alpha = attr_f(node, "alpha", 1.0) as f64;
1369 let beta = attr_f(node, "beta", 1.0) as f64;
1370 let trans_a = attr_i(node, "transA", 0) != 0;
1371 let trans_b = attr_i(node, "transB", 0) != 0;
1372
1373 let mut a = get_tensor(env, &node.inputs[0])?.clone();
1374 let mut b = get_tensor(env, &node.inputs[1])?.clone();
1375
1376 if trans_a {
1377 a = a.t()?;
1378 }
1379 if trans_b {
1380 b = b.t()?;
1381 }
1382
1383 let mut out = a.matmul(&b)?;
1384 if (alpha - 1.0).abs() > 1e-7 {
1385 out = out.affine(alpha, 0.0)?;
1386 }
1387 if node.inputs.len() > 2 && !node.inputs[2].is_empty() {
1388 let c = get_tensor(env, &node.inputs[2])?;
1389 if (beta - 1.0).abs() > 1e-7 {
1390 let bc = c.affine(beta, 0.0)?;
1391 out = out.add(&bc)?;
1392 } else {
1393 out = out.add(c)?;
1394 }
1395 }
1396 env.insert(node.outputs[0].clone(), out);
1397 }
1398
1399 "Relu" => {
1401 let x = get_tensor(env, &node.inputs[0])?;
1402 env.insert(node.outputs[0].clone(), x.relu()?);
1403 }
1404 "Sigmoid" => {
1405 let x = get_tensor(env, &node.inputs[0])?;
1406 env.insert(node.outputs[0].clone(), x.sigmoid()?);
1407 }
1408 "Tanh" => {
1409 let x = get_tensor(env, &node.inputs[0])?;
1410 env.insert(node.outputs[0].clone(), x.tanh()?);
1411 }
1412 "Neg" => {
1413 let x = get_tensor(env, &node.inputs[0])?;
1414 env.insert(node.outputs[0].clone(), x.neg()?);
1415 }
1416 "Sqrt" => {
1417 let x = get_tensor(env, &node.inputs[0])?;
1418 env.insert(node.outputs[0].clone(), x.sqrt()?);
1419 }
1420 "Exp" => {
1421 let x = get_tensor(env, &node.inputs[0])?;
1422 env.insert(node.outputs[0].clone(), x.exp()?);
1423 }
1424 "Log" => {
1425 let x = get_tensor(env, &node.inputs[0])?;
1426 env.insert(node.outputs[0].clone(), x.log()?);
1427 }
1428 "Abs" => {
1429 let x = get_tensor(env, &node.inputs[0])?;
1430 env.insert(node.outputs[0].clone(), x.abs()?);
1431 }
1432
1433 "Softmax" => {
1435 let x = get_tensor(env, &node.inputs[0])?;
1436 let axis = attr_i(node, "axis", -1);
1437 let dim = if axis < 0 {
1438 (x.rank() as i64 + axis) as usize
1439 } else {
1440 axis as usize
1441 };
1442 env.insert(node.outputs[0].clone(), x.softmax(dim)?);
1443 }
1444 "LogSoftmax" => {
1445 let x = get_tensor(env, &node.inputs[0])?;
1446 let axis = attr_i(node, "axis", -1);
1447 let dim = if axis < 0 {
1448 (x.rank() as i64 + axis) as usize
1449 } else {
1450 axis as usize
1451 };
1452 env.insert(node.outputs[0].clone(), x.log_softmax(dim)?);
1453 }
1454
1455 "Clip" => {
1457 let x = get_tensor(env, &node.inputs[0])?;
1458 let min_val = if node.inputs.len() > 1 && !node.inputs[1].is_empty() {
1459 get_tensor(env, &node.inputs[1])?.to_f64_vec()?[0]
1460 } else {
1461 f64::NEG_INFINITY
1462 };
1463 let max_val = if node.inputs.len() > 2 && !node.inputs[2].is_empty() {
1464 get_tensor(env, &node.inputs[2])?.to_f64_vec()?[0]
1465 } else {
1466 f64::INFINITY
1467 };
1468 env.insert(node.outputs[0].clone(), x.clamp(min_val, max_val)?);
1469 }
1470
1471 "Reshape" => {
1473 let x = get_tensor(env, &node.inputs[0])?;
1474 let shape_tensor = get_tensor(env, &node.inputs[1])?;
1475 let shape_vals = shape_tensor.to_f64_vec()?;
1476
1477 let total = x.elem_count();
1479 let mut new_shape: Vec<usize> = shape_vals.iter().map(|&v| v as i64 as usize).collect();
1480 let neg_idx = new_shape.iter().position(|&s| s == usize::MAX); if let Some(idx) = neg_idx {
1482 let known: usize = new_shape
1483 .iter()
1484 .enumerate()
1485 .filter(|&(i, _)| i != idx)
1486 .map(|(_, &s)| s)
1487 .product();
1488 if known > 0 {
1489 new_shape[idx] = total / known;
1490 }
1491 }
1492 env.insert(node.outputs[0].clone(), x.reshape(new_shape)?);
1493 }
1494 "Transpose" => {
1495 let x = get_tensor(env, &node.inputs[0])?;
1496 let perm = attr_ints(node, "perm");
1497 if perm.is_empty() {
1498 let rank = x.rank();
1500 let rev: Vec<usize> = (0..rank).rev().collect();
1501 env.insert(node.outputs[0].clone(), x.permute(&rev)?);
1502 } else {
1503 let perm_usize: Vec<usize> = perm.iter().map(|&p| p as usize).collect();
1504 env.insert(node.outputs[0].clone(), x.permute(&perm_usize)?);
1505 }
1506 }
1507 "Flatten" => {
1508 let x = get_tensor(env, &node.inputs[0])?;
1509 let axis = attr_i(node, "axis", 1) as usize;
1510 env.insert(node.outputs[0].clone(), x.flatten(axis, x.rank() - 1)?);
1511 }
1512 "Squeeze" => {
1513 let x = get_tensor(env, &node.inputs[0])?;
1514 let axes = attr_ints(node, "axes");
1515 if axes.is_empty() {
1516 env.insert(node.outputs[0].clone(), x.squeeze_all());
1517 } else {
1518 let mut result = x.clone();
1519 let mut sorted_axes: Vec<usize> = axes.iter().map(|&a| a as usize).collect();
1521 sorted_axes.sort_unstable();
1522 sorted_axes.reverse();
1523 for ax in sorted_axes {
1524 result = result.squeeze(ax)?;
1525 }
1526 env.insert(node.outputs[0].clone(), result);
1527 }
1528 }
1529 "Unsqueeze" => {
1530 let x = get_tensor(env, &node.inputs[0])?;
1531 let axes = if node.inputs.len() > 1 && !node.inputs[1].is_empty() {
1532 let axes_t = get_tensor(env, &node.inputs[1])?;
1534 axes_t
1535 .to_f64_vec()?
1536 .iter()
1537 .map(|&v| v as i64)
1538 .collect::<Vec<_>>()
1539 } else {
1540 attr_ints(node, "axes")
1541 };
1542 let mut result = x.clone();
1543 let mut sorted_axes: Vec<usize> = axes
1544 .iter()
1545 .map(|&a| {
1546 if a < 0 {
1547 (result.rank() as i64 + a + 1) as usize
1548 } else {
1549 a as usize
1550 }
1551 })
1552 .collect();
1553 sorted_axes.sort_unstable();
1554 for ax in sorted_axes {
1555 result = result.unsqueeze(ax)?;
1556 }
1557 env.insert(node.outputs[0].clone(), result);
1558 }
1559 "Concat" => {
1560 let axis = attr_i(node, "axis", 0) as usize;
1561 let tensors: Vec<Tensor<B>> = node
1562 .inputs
1563 .iter()
1564 .map(|n| get_tensor(env, n).cloned())
1565 .collect::<Result<Vec<_>>>()?;
1566 let refs: Vec<Tensor<B>> = tensors;
1567 let out = Tensor::<B>::cat(&refs, axis)?;
1568 env.insert(node.outputs[0].clone(), out);
1569 }
1570
1571 "ReduceSum" => {
1573 let x = get_tensor(env, &node.inputs[0])?;
1574 let axes = attr_ints(node, "axes");
1575 let keepdims = attr_i(node, "keepdims", 1) != 0;
1576 let mut result = x.clone();
1577 if axes.is_empty() {
1578 result = result.sum_all()?;
1579 } else {
1580 let mut sorted: Vec<usize> = axes.iter().map(|&a| a as usize).collect();
1581 sorted.sort_unstable();
1582 sorted.reverse();
1583 for ax in sorted {
1584 result = result.sum(ax, keepdims)?;
1585 }
1586 }
1587 env.insert(node.outputs[0].clone(), result);
1588 }
1589 "ReduceMean" => {
1590 let x = get_tensor(env, &node.inputs[0])?;
1591 let axes = attr_ints(node, "axes");
1592 let keepdims = attr_i(node, "keepdims", 1) != 0;
1593 let mut result = x.clone();
1594 if axes.is_empty() {
1595 result = result.mean_all()?;
1596 } else {
1597 let mut sorted: Vec<usize> = axes.iter().map(|&a| a as usize).collect();
1598 sorted.sort_unstable();
1599 sorted.reverse();
1600 for ax in sorted {
1601 result = result.mean(ax, keepdims)?;
1602 }
1603 }
1604 env.insert(node.outputs[0].clone(), result);
1605 }
1606 "ReduceMax" => {
1607 let x = get_tensor(env, &node.inputs[0])?;
1608 let axes = attr_ints(node, "axes");
1609 let keepdims = attr_i(node, "keepdims", 1) != 0;
1610 let mut result = x.clone();
1611 let mut sorted: Vec<usize> = axes.iter().map(|&a| a as usize).collect();
1612 sorted.sort_unstable();
1613 sorted.reverse();
1614 for ax in sorted {
1615 result = result.max(ax, keepdims)?;
1616 }
1617 env.insert(node.outputs[0].clone(), result);
1618 }
1619 "ReduceMin" => {
1620 let x = get_tensor(env, &node.inputs[0])?;
1621 let axes = attr_ints(node, "axes");
1622 let keepdims = attr_i(node, "keepdims", 1) != 0;
1623 let mut result = x.clone();
1624 let mut sorted: Vec<usize> = axes.iter().map(|&a| a as usize).collect();
1625 sorted.sort_unstable();
1626 sorted.reverse();
1627 for ax in sorted {
1628 result = result.min(ax, keepdims)?;
1629 }
1630 env.insert(node.outputs[0].clone(), result);
1631 }
1632
1633 "Gather" => {
1635 let x = get_tensor(env, &node.inputs[0])?;
1636 let indices = get_tensor(env, &node.inputs[1])?;
1637 let axis = attr_i(node, "axis", 0) as usize;
1638 env.insert(node.outputs[0].clone(), x.gather(axis, indices)?);
1639 }
1640
1641 "BatchNormalization" => {
1643 let x = get_tensor(env, &node.inputs[0])?;
1645 let scale = get_tensor(env, &node.inputs[1])?;
1646 let bias = get_tensor(env, &node.inputs[2])?;
1647 let mean = get_tensor(env, &node.inputs[3])?;
1648 let var = get_tensor(env, &node.inputs[4])?;
1649 let eps = attr_f(node, "epsilon", 1e-5) as f64;
1650
1651 let x_sub = x.sub(mean)?;
1654 let std_inv = var.affine(1.0, eps)?.sqrt()?.reciprocal()?;
1655 let normed = x_sub.mul(&std_inv)?;
1656 let scaled = normed.mul(scale)?;
1657 let out = scaled.add(bias)?;
1658 env.insert(node.outputs[0].clone(), out);
1659 }
1660
1661 "Dropout" => {
1663 let x = get_tensor(env, &node.inputs[0])?.clone();
1665 env.insert(node.outputs[0].clone(), x.clone());
1666 if node.outputs.len() > 1 && !node.outputs[1].is_empty() {
1668 env.insert(node.outputs[1].clone(), x);
1669 }
1670 }
1671
1672 "Identity" => {
1674 let x = get_tensor(env, &node.inputs[0])?;
1675 env.insert(node.outputs[0].clone(), x.clone());
1676 }
1677
1678 "Shape" => {
1680 let x = get_tensor(env, &node.inputs[0])?;
1681 let shape: Vec<f64> = x.dims().iter().map(|&d| d as f64).collect();
1682 let n = shape.len();
1683 let out = Tensor::<B>::from_f64_slice(&shape, vec![n], DType::I64, device)?;
1684 env.insert(node.outputs[0].clone(), out);
1685 }
1686
1687 "Cast" => {
1689 let x = get_tensor(env, &node.inputs[0])?;
1690 let to = attr_i(node, "to", ONNX_FLOAT as i64) as i32;
1691 let target_dtype = onnx_to_dtype(to)?;
1692 env.insert(node.outputs[0].clone(), x.to_dtype(target_dtype)?);
1693 }
1694
1695 "Constant" => {
1697 if let Some(OnnxAttribute::Float(v)) = node.attributes.get("value_float") {
1699 let out = Tensor::<B>::from_f64_slice(&[*v as f64], vec![1], DType::F32, device)?;
1700 env.insert(node.outputs[0].clone(), out);
1701 } else if let Some(OnnxAttribute::Int(v)) = node.attributes.get("value_int") {
1702 let out = Tensor::<B>::from_f64_slice(&[*v as f64], vec![1], DType::I64, device)?;
1703 env.insert(node.outputs[0].clone(), out);
1704 } else if let Some(OnnxAttribute::Floats(v)) = node.attributes.get("value_floats") {
1705 let data: Vec<f64> = v.iter().map(|f| *f as f64).collect();
1706 let n = data.len();
1707 let out = Tensor::<B>::from_f64_slice(&data, vec![n], DType::F32, device)?;
1708 env.insert(node.outputs[0].clone(), out);
1709 } else if let Some(OnnxAttribute::Ints(v)) = node.attributes.get("value_ints") {
1710 let data: Vec<f64> = v.iter().map(|i| *i as f64).collect();
1711 let n = data.len();
1712 let out = Tensor::<B>::from_f64_slice(&data, vec![n], DType::I64, device)?;
1713 env.insert(node.outputs[0].clone(), out);
1714 } else {
1715 return Err(shrew_core::Error::msg(format!(
1716 "ONNX Constant: unsupported value attribute in node '{}'",
1717 node.name
1718 )));
1719 }
1720 }
1721
1722 other => {
1723 return Err(shrew_core::Error::msg(format!(
1724 "ONNX runtime: unsupported op '{other}' (node '{}')",
1725 node.name
1726 )));
1727 }
1728 }
1729 Ok(())
1730}
1731
1732#[cfg(test)]
1737mod tests {
1738 use super::*;
1739 use shrew_cpu::{CpuBackend, CpuDevice};
1740
1741 type B = CpuBackend;
1742 type T = Tensor<B>;
1743 const DEV: CpuDevice = CpuDevice;
1744
1745 #[test]
1746 fn test_protobuf_varint_roundtrip() {
1747 let mut enc = PbEncoder::new();
1748 enc.write_varint(0);
1749 enc.write_varint(1);
1750 enc.write_varint(127);
1751 enc.write_varint(128);
1752 enc.write_varint(300);
1753 enc.write_varint(16384);
1754
1755 let mut dec = PbDecoder::new(&enc.buf);
1756 assert_eq!(dec.read_varint().unwrap(), 0);
1757 assert_eq!(dec.read_varint().unwrap(), 1);
1758 assert_eq!(dec.read_varint().unwrap(), 127);
1759 assert_eq!(dec.read_varint().unwrap(), 128);
1760 assert_eq!(dec.read_varint().unwrap(), 300);
1761 assert_eq!(dec.read_varint().unwrap(), 16384);
1762 }
1763
1764 #[test]
1765 fn test_onnx_tensor_encode_decode() {
1766 let mut tensor = OnnxTensor::new("test_weight");
1767 tensor.data_type = ONNX_FLOAT;
1768 tensor.dims = vec![2, 3];
1769 tensor.float_data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1770
1771 let encoded = tensor.encode();
1772 let decoded = OnnxTensor::decode(&encoded).unwrap();
1773
1774 assert_eq!(decoded.name, "test_weight");
1775 assert_eq!(decoded.data_type, ONNX_FLOAT);
1776 assert_eq!(decoded.dims, vec![2, 3]);
1777 let data = decoded.to_f64_vec();
1779 assert_eq!(data.len(), 6);
1780 for (a, b) in tensor.float_data.iter().zip(data.iter()) {
1781 assert!((*a as f64 - b).abs() < 1e-6);
1782 }
1783 }
1784
1785 #[test]
1786 fn test_export_import_roundtrip() {
1787 let linear = shrew_nn::Linear::<B>::new(4, 3, true, DType::F32, &DEV).unwrap();
1788
1789 let path = std::env::temp_dir().join("shrew_test_onnx.onnx");
1791 export_weights(&path, &linear, "test_model", &[1, 4]).unwrap();
1792
1793 let weights = load_onnx_weights::<B>(&path, &DEV).unwrap();
1795
1796 assert_eq!(weights.len(), 2);
1798
1799 let named = linear.named_parameters();
1801 for (name, original) in &named {
1802 let loaded = weights.get(name).expect(&format!("missing: {name}"));
1803 assert_eq!(original.dims(), loaded.dims());
1804
1805 let orig_data = original.to_f64_vec().unwrap();
1807 let load_data = loaded.to_f64_vec().unwrap();
1808 for (a, b) in orig_data.iter().zip(load_data.iter()) {
1809 assert!((a - b).abs() < 1e-5, "mismatch for {name}: {a} vs {b}");
1810 }
1811 }
1812
1813 let _ = fs::remove_file(&path);
1815 }
1816
1817 #[test]
1818 fn test_export_tensors() {
1819 let t1 = T::randn(vec![2, 3], DType::F32, &DEV).unwrap();
1820 let t2 = T::ones(vec![5], DType::F32, &DEV).unwrap();
1821
1822 let tensors = vec![
1823 ("weight".to_string(), t1.clone()),
1824 ("bias".to_string(), t2.clone()),
1825 ];
1826
1827 let path = std::env::temp_dir().join("shrew_test_tensors.onnx");
1828 export_tensors(&path, &tensors, "tensors_model").unwrap();
1829
1830 let loaded = load_onnx_weights::<B>(&path, &DEV).unwrap();
1831 assert_eq!(loaded.len(), 2);
1832 assert_eq!(loaded.get("weight").unwrap().dims(), &[2, 3]);
1833 assert_eq!(loaded.get("bias").unwrap().dims(), &[5]);
1834
1835 let _ = fs::remove_file(&path);
1836 }
1837
1838 #[test]
1839 fn test_onnx_model_builder() {
1840 let mut model = OnnxModel::new("test_graph");
1841 model
1842 .inputs
1843 .push(("X".to_string(), vec![1, 784], ONNX_FLOAT));
1844 model
1845 .outputs
1846 .push(("Y".to_string(), vec![1, 10], ONNX_FLOAT));
1847
1848 model.nodes.push(OnnxNode {
1849 inputs: vec!["X".to_string(), "weight".to_string()],
1850 outputs: vec!["matmul_out".to_string()],
1851 op_type: "MatMul".to_string(),
1852 name: "matmul_0".to_string(),
1853 attributes: HashMap::new(),
1854 });
1855
1856 let mut attrs = HashMap::new();
1857 attrs.insert("axis".to_string(), OnnxAttribute::Int(1));
1858 model.nodes.push(OnnxNode {
1859 inputs: vec!["matmul_out".to_string()],
1860 outputs: vec!["Y".to_string()],
1861 op_type: "Softmax".to_string(),
1862 name: "softmax_0".to_string(),
1863 attributes: attrs,
1864 });
1865
1866 let bytes = model.to_bytes();
1867 assert!(!bytes.is_empty());
1868 assert!(bytes.len() > 20); }
1870
1871 #[test]
1872 fn test_dtype_conversion() {
1873 assert_eq!(dtype_to_onnx(DType::F32), ONNX_FLOAT);
1874 assert_eq!(dtype_to_onnx(DType::F64), ONNX_DOUBLE);
1875 assert_eq!(dtype_to_onnx(DType::F16), ONNX_FLOAT16);
1876 assert_eq!(onnx_to_dtype(ONNX_FLOAT).unwrap(), DType::F32);
1877 assert_eq!(onnx_to_dtype(ONNX_DOUBLE).unwrap(), DType::F64);
1878 assert_eq!(onnx_to_dtype(ONNX_FLOAT16).unwrap(), DType::F16);
1879 }
1880
1881 #[test]
1882 fn test_double_data_roundtrip() {
1883 let t = T::from_f64_slice(&[1.0, 2.0, 3.0, 4.0], vec![2, 2], DType::F64, &DEV).unwrap();
1884
1885 let tensors = vec![("w".to_string(), t.clone())];
1886 let path = std::env::temp_dir().join("shrew_test_f64.onnx");
1887 export_tensors(&path, &tensors, "f64_model").unwrap();
1888
1889 let loaded = load_onnx_weights::<B>(&path, &DEV).unwrap();
1890 let w = loaded.get("w").unwrap();
1891 assert_eq!(w.dims(), &[2, 2]);
1892
1893 let orig = t.to_f64_vec().unwrap();
1894 let load = w.to_f64_vec().unwrap();
1895 for (a, b) in orig.iter().zip(load.iter()) {
1896 assert!((a - b).abs() < 1e-10);
1897 }
1898
1899 let _ = fs::remove_file(&path);
1900 }
1901
1902 fn build_and_reload_graph(model: &OnnxModel) -> OnnxGraph {
1908 let bytes = model.to_bytes();
1909 load_onnx_graph_from_bytes(&bytes).unwrap()
1910 }
1911
1912 #[test]
1913 fn test_graph_add_two_inputs() {
1914 let mut model = OnnxModel::new("add_graph");
1916 model.inputs.push(("A".into(), vec![2, 2], ONNX_FLOAT));
1917 model.inputs.push(("B".into(), vec![2, 2], ONNX_FLOAT));
1918 model.outputs.push(("Y".into(), vec![2, 2], ONNX_FLOAT));
1919 model.nodes.push(OnnxNode {
1920 inputs: vec!["A".into(), "B".into()],
1921 outputs: vec!["Y".into()],
1922 op_type: "Add".into(),
1923 name: "add_0".into(),
1924 attributes: HashMap::new(),
1925 });
1926
1927 let graph = build_and_reload_graph(&model);
1928 assert_eq!(graph.nodes.len(), 1);
1929 assert_eq!(graph.nodes[0].op_type, "Add");
1930 assert_eq!(graph.output_names, vec!["Y"]);
1931
1932 let a = T::from_f64_slice(&[1.0, 2.0, 3.0, 4.0], vec![2, 2], DType::F32, &DEV).unwrap();
1933 let b = T::from_f64_slice(&[10.0, 20.0, 30.0, 40.0], vec![2, 2], DType::F32, &DEV).unwrap();
1934
1935 let mut inputs = HashMap::new();
1936 inputs.insert("A".into(), a);
1937 inputs.insert("B".into(), b);
1938
1939 let outputs = run_onnx_graph::<B>(&graph, &inputs, &DEV).unwrap();
1940 let y = outputs.get("Y").unwrap();
1941 let data = y.to_f64_vec().unwrap();
1942 assert_eq!(data, vec![11.0, 22.0, 33.0, 44.0]);
1943 }
1944
1945 #[test]
1946 fn test_graph_linear_relu() {
1947 let mut model = OnnxModel::new("linear_relu");
1952 model.inputs.push(("X".into(), vec![1, 2], ONNX_FLOAT));
1953 model.outputs.push(("Z".into(), vec![1, 3], ONNX_FLOAT));
1954
1955 let mut w = OnnxTensor::new("W");
1957 w.data_type = ONNX_FLOAT;
1958 w.dims = vec![2, 3];
1959 w.float_data = vec![1.0, -1.0, 0.5, 0.0, 2.0, -0.5];
1960 model.initializers.push(w);
1961
1962 let mut b = OnnxTensor::new("B");
1963 b.data_type = ONNX_FLOAT;
1964 b.dims = vec![3];
1965 b.float_data = vec![0.0, 0.0, 0.0];
1966 model.initializers.push(b);
1967
1968 model.nodes.push(OnnxNode {
1969 inputs: vec!["X".into(), "W".into()],
1970 outputs: vec!["matmul_out".into()],
1971 op_type: "MatMul".into(),
1972 name: "matmul_0".into(),
1973 attributes: HashMap::new(),
1974 });
1975 model.nodes.push(OnnxNode {
1976 inputs: vec!["matmul_out".into(), "B".into()],
1977 outputs: vec!["add_out".into()],
1978 op_type: "Add".into(),
1979 name: "add_0".into(),
1980 attributes: HashMap::new(),
1981 });
1982 model.nodes.push(OnnxNode {
1983 inputs: vec!["add_out".into()],
1984 outputs: vec!["Z".into()],
1985 op_type: "Relu".into(),
1986 name: "relu_0".into(),
1987 attributes: HashMap::new(),
1988 });
1989
1990 let graph = build_and_reload_graph(&model);
1991 assert_eq!(graph.nodes.len(), 3);
1992
1993 let x = T::from_f64_slice(&[1.0, -1.0], vec![1, 2], DType::F32, &DEV).unwrap();
1998 let mut inputs = HashMap::new();
1999 inputs.insert("X".into(), x);
2000
2001 let outputs = run_onnx_graph::<B>(&graph, &inputs, &DEV).unwrap();
2002 let z = outputs.get("Z").unwrap();
2003 let data = z.to_f64_vec().unwrap();
2004 assert_eq!(data.len(), 3);
2005 assert!((data[0] - 1.0).abs() < 1e-5);
2006 assert!((data[1] - 0.0).abs() < 1e-5);
2007 assert!((data[2] - 1.0).abs() < 1e-5);
2008 }
2009
2010 #[test]
2011 fn test_graph_identity_and_dropout() {
2012 let mut model = OnnxModel::new("id_drop");
2013 model.inputs.push(("X".into(), vec![3], ONNX_FLOAT));
2014 model.outputs.push(("Y".into(), vec![3], ONNX_FLOAT));
2015
2016 model.nodes.push(OnnxNode {
2017 inputs: vec!["X".into()],
2018 outputs: vec!["id_out".into()],
2019 op_type: "Identity".into(),
2020 name: "id_0".into(),
2021 attributes: HashMap::new(),
2022 });
2023 model.nodes.push(OnnxNode {
2024 inputs: vec!["id_out".into()],
2025 outputs: vec!["Y".into()],
2026 op_type: "Dropout".into(),
2027 name: "drop_0".into(),
2028 attributes: HashMap::new(),
2029 });
2030
2031 let graph = build_and_reload_graph(&model);
2032 let x = T::from_f64_slice(&[5.0, -3.0, 7.0], vec![3], DType::F32, &DEV).unwrap();
2033 let mut inputs = HashMap::new();
2034 inputs.insert("X".into(), x.clone());
2035
2036 let outputs = run_onnx_graph::<B>(&graph, &inputs, &DEV).unwrap();
2037 let y = outputs.get("Y").unwrap();
2038 assert_eq!(y.to_f64_vec().unwrap(), x.to_f64_vec().unwrap());
2039 }
2040
2041 #[test]
2042 fn test_graph_file_roundtrip() {
2043 let mut model = OnnxModel::new("file_rt");
2045 model.inputs.push(("X".into(), vec![2], ONNX_FLOAT));
2046 model.outputs.push(("Y".into(), vec![2], ONNX_FLOAT));
2047
2048 model.nodes.push(OnnxNode {
2049 inputs: vec!["X".into()],
2050 outputs: vec!["Y".into()],
2051 op_type: "Sigmoid".into(),
2052 name: "sig_0".into(),
2053 attributes: HashMap::new(),
2054 });
2055
2056 let path = std::env::temp_dir().join("shrew_test_graph_rt.onnx");
2057 model.save(&path).unwrap();
2058
2059 let graph = load_onnx_graph(&path).unwrap();
2060 assert_eq!(graph.nodes.len(), 1);
2061
2062 let x = T::from_f64_slice(&[0.0, 1000.0], vec![2], DType::F32, &DEV).unwrap();
2063 let mut inputs = HashMap::new();
2064 inputs.insert("X".into(), x);
2065
2066 let outputs = run_onnx_graph::<B>(&graph, &inputs, &DEV).unwrap();
2067 let data = outputs.get("Y").unwrap().to_f64_vec().unwrap();
2068 assert!((data[0] - 0.5).abs() < 1e-5); assert!((data[1] - 1.0).abs() < 1e-3); let _ = fs::remove_file(&path);
2072 }
2073
2074 #[test]
2075 fn test_decode_attribute_roundtrip() {
2076 let attr = OnnxAttribute::Int(42);
2078 let encoded = encode_attribute("axis", &attr);
2079 let (name, decoded) = decode_attribute(&encoded.buf).unwrap();
2080 assert_eq!(name, "axis");
2081 match decoded {
2082 OnnxAttribute::Int(v) => assert_eq!(v, 42),
2083 _ => panic!("expected Int"),
2084 }
2085 }
2086
2087 #[test]
2088 fn test_ir_graph_to_onnx_roundtrip() {
2089 let mut ir_graph = IrGraph::new("demo_ir");
2090 let x = ir_graph.add_node(
2091 "x",
2092 OpKind::Identity,
2093 vec![],
2094 IrType::Tensor {
2095 shape: vec![Dim::Fixed(2), Dim::Fixed(2)],
2096 dtype: IrDType::F32,
2097 },
2098 );
2099 ir_graph.inputs.push(x);
2100
2101 let relu = ir_graph.add_node(
2102 "out",
2103 OpKind::Relu,
2104 vec![x],
2105 IrType::Tensor {
2106 shape: vec![Dim::Fixed(2), Dim::Fixed(2)],
2107 dtype: IrDType::F32,
2108 },
2109 );
2110 ir_graph.add_output(relu);
2111
2112 let path = std::env::temp_dir().join("shrew_test_ir_export.onnx");
2113 export_ir_graph(&path, &ir_graph).unwrap();
2114
2115 let loaded = load_onnx_graph(&path).unwrap();
2116 assert_eq!(loaded.nodes.len(), 1);
2117 assert_eq!(loaded.nodes[0].op_type, "Relu");
2118 assert_eq!(loaded.output_names, vec!["out"]);
2119
2120 let t = T::from_f64_slice(&[-2.0, 3.0, -1.0, 4.0], vec![2, 2], DType::F32, &DEV).unwrap();
2121 let mut inputs = HashMap::new();
2122 inputs.insert("x".into(), t);
2123
2124 let outputs = run_onnx_graph::<B>(&loaded, &inputs, &DEV).unwrap();
2125 let out = outputs.get("out").unwrap();
2126 assert_eq!(out.to_f64_vec().unwrap(), vec![0.0, 3.0, 0.0, 4.0]);
2127
2128 let _ = fs::remove_file(&path);
2129 }
2130}