Skip to main content

shrew/
onnx.rs

1// =============================================================================
2// ONNX — Import / Export for interoperability
3// =============================================================================
4//
5// ONNX (Open Neural Network Exchange) is the industry standard for
6// exchanging trained models between frameworks (PyTorch, TensorFlow,
7// CoreML, TensorRT, etc.).
8//
9// This module provides:
10//
11//   - Export: convert a Shrew module's state_dict to ONNX format
12//   - Import: load an ONNX model's weights into Shrew tensors
13//
14// ONNX files use Protocol Buffers encoding. We implement a minimal
15// protobuf encoder/decoder (no external crate needed) that handles
16// the subset of the ONNX spec we need: ModelProto, GraphProto,
17// TensorProto, and NodeProto.
18//
19// SUPPORTED ONNX OPS (for graph export):
20//   MatMul, Add, Relu, Sigmoid, Tanh, Softmax, Gemm, Reshape, Transpose,
21//   Conv, BatchNormalization, Dropout, Concat, Flatten
22//
23// REFERENCE:
24//   https://onnx.ai/onnx/repo-docs/IR.html
25//   https://protobuf.dev/programming-guides/encoding/
26
27use 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
39// =============================================================================
40// ONNX constants
41// =============================================================================
42
43/// ONNX IR version (we target ONNX IR version 9 / opset 17).
44const ONNX_IR_VERSION: i64 = 9;
45/// Default opset version.
46const ONNX_OPSET_VERSION: i64 = 17;
47/// Magic bytes + version for ONNX protobuf.
48const ONNX_DOMAIN: &str = "";
49
50// ONNX TensorProto data types
51/// See https://onnx.ai/onnx/repo-docs/IR.html#tensor-data-types
52const 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
62// =============================================================================
63// Minimal protobuf encoder
64// =============================================================================
65
66/// A minimal protobuf wire-format encoder. Supports:
67/// - Varint (field type 0)
68/// - Length-delimited (field type 2: bytes, strings, nested messages)
69/// - Fixed32/Fixed64 (field types 5 and 1)
70struct 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    /// Write a varint.
84    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    /// Write a field tag (field_number << 3 | wire_type).
98    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    /// Write a varint field.
103    fn write_varint_field(&mut self, field: u32, val: u64) {
104        self.write_tag(field, 0);
105        self.write_varint(val);
106    }
107
108    /// Write a signed varint field (zigzag encoding for negative values).
109    fn write_sint64_field(&mut self, field: u32, val: i64) {
110        self.write_varint_field(field, val as u64);
111    }
112
113    /// Write a length-delimited bytes field.
114    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    /// Write a string field.
121    fn write_string_field(&mut self, field: u32, val: &str) {
122        self.write_bytes_field(field, val.as_bytes());
123    }
124
125    /// Write a nested message field.
126    fn write_message_field(&mut self, field: u32, encoder: &PbEncoder) {
127        self.write_bytes_field(field, &encoder.buf);
128    }
129
130    /// Write raw float data as bytes.
131    #[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    /// Write raw double data as bytes.
138    #[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
145// =============================================================================
146// Minimal protobuf decoder
147// =============================================================================
148
149/// A minimal protobuf wire-format decoder.
150struct 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            } // fixed64
216            2 => {
217                self.read_bytes()?;
218            }
219            5 => {
220                self.pos += 4;
221            } // fixed32
222            _ => {
223                return Err(shrew_core::Error::msg(format!(
224                    "protobuf: unsupported wire type {wire_type}"
225                )))
226            }
227        }
228        Ok(())
229    }
230}
231
232// =============================================================================
233// ONNX TensorProto
234// =============================================================================
235
236/// Represents an ONNX TensorProto (a named tensor with shape and data).
237#[derive(Debug, Clone)]
238pub struct OnnxTensor {
239    /// Tensor name.
240    pub name: String,
241    /// ONNX data type (ONNX_FLOAT, ONNX_DOUBLE, etc.).
242    pub data_type: i32,
243    /// Shape dimensions.
244    pub dims: Vec<i64>,
245    /// Raw float data (for FLOAT type).
246    pub float_data: Vec<f32>,
247    /// Raw double data (for DOUBLE type).
248    pub double_data: Vec<f64>,
249    /// Raw bytes (for packed formats like FLOAT16).
250    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    /// Convert to protobuf bytes.
266    fn encode(&self) -> Vec<u8> {
267        let mut enc = PbEncoder::new();
268        // field 1: dims (repeated int64)
269        for &d in &self.dims {
270            enc.write_sint64_field(1, d);
271        }
272        // field 2: data_type (int32)
273        enc.write_varint_field(2, self.data_type as u64);
274        // field 8: name (string)
275        if !self.name.is_empty() {
276            enc.write_string_field(8, &self.name);
277        }
278        // field 4: float_data (packed repeated float — as raw_data for efficiency)
279        if !self.float_data.is_empty() {
280            // field 13: raw_data (bytes) — more efficient than repeated float
281            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    /// Decode from protobuf bytes.
301    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                    // dims (varint)
309                    let v = dec.read_varint()? as i64;
310                    tensor.dims.push(v);
311                }
312                (1, 2) => {
313                    // dims (packed)
314                    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                    // data_type
322                    tensor.data_type = dec.read_varint()? as i32;
323                }
324                (8, 2) => {
325                    // name
326                    tensor.name = dec.read_string()?;
327                }
328                (13, 2) => {
329                    // raw_data
330                    tensor.raw_data = dec.read_bytes()?.to_vec();
331                }
332                (4, 2) => {
333                    // float_data (packed)
334                    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                    // float_data (repeated fixed32)
342                    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                    // double_data (packed)
349                    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    /// Get the float data (converting from raw_data if needed).
367    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
402/// Convert a Shrew DType to ONNX data type integer.
403fn 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
415/// Convert an ONNX data type integer to Shrew DType.
416fn 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),   // map to u8
426        ONNX_INT32 => Ok(DType::I64), // upcast
427        _ => Err(shrew_core::Error::msg(format!(
428            "unsupported ONNX data type: {onnx_type}"
429        ))),
430    }
431}
432
433// =============================================================================
434// ONNX NodeProto (graph operation)
435// =============================================================================
436
437/// An ONNX graph node (operation).
438#[derive(Debug, Clone)]
439pub struct OnnxNode {
440    /// Input tensor names.
441    pub inputs: Vec<String>,
442    /// Output tensor names.
443    pub outputs: Vec<String>,
444    /// Operation type (e.g., "MatMul", "Relu", "Add").
445    pub op_type: String,
446    /// Node name (for debugging).
447    pub name: String,
448    /// String attributes (key → value).
449    pub attributes: HashMap<String, OnnxAttribute>,
450}
451
452/// An ONNX attribute value.
453#[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        // field 1: inputs (repeated string)
466        for input in &self.inputs {
467            enc.write_string_field(1, input);
468        }
469        // field 2: outputs (repeated string)
470        for output in &self.outputs {
471            enc.write_string_field(2, output);
472        }
473        // field 3: name (string)
474        if !self.name.is_empty() {
475            enc.write_string_field(3, &self.name);
476        }
477        // field 4: op_type (string)
478        enc.write_string_field(4, &self.op_type);
479        // field 5: attributes (repeated AttributeProto)
480        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); // field 1: name
491    match val {
492        OnnxAttribute::Int(i) => {
493            enc.write_varint_field(2, 2); // type = INT
494            enc.write_sint64_field(3, *i); // field 3: i
495        }
496        OnnxAttribute::Float(f) => {
497            enc.write_varint_field(2, 1); // type = FLOAT
498                                          // field 4: f (float, fixed32)
499            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); // type = STRING
504            enc.write_bytes_field(5, s.as_bytes()); // field 5: s
505        }
506        OnnxAttribute::Ints(ints) => {
507            enc.write_varint_field(2, 7); // type = INTS
508            for &i in ints {
509                enc.write_sint64_field(8, i); // field 8: ints
510            }
511        }
512        OnnxAttribute::Floats(floats) => {
513            enc.write_varint_field(2, 6); // type = FLOATS
514            for &f in floats {
515                enc.write_tag(7, 5); // field 7: floats (fixed32)
516                enc.buf.extend_from_slice(&f.to_le_bytes());
517            }
518        }
519    }
520    enc
521}
522
523// =============================================================================
524// ONNX ModelProto — top-level export
525// =============================================================================
526
527/// An ONNX model with graph, metadata, and opset information.
528#[derive(Debug, Clone)]
529pub struct OnnxModel {
530    /// Model producer name.
531    pub producer_name: String,
532    /// Model producer version.
533    pub producer_version: String,
534    /// Graph name.
535    pub graph_name: String,
536    /// Graph nodes (operations).
537    pub nodes: Vec<OnnxNode>,
538    /// Initializer tensors (weights).
539    pub initializers: Vec<OnnxTensor>,
540    /// Graph inputs (names and shapes).
541    pub inputs: Vec<(String, Vec<i64>, i32)>,
542    /// Graph outputs (names and shapes).
543    pub outputs: Vec<(String, Vec<i64>, i32)>,
544}
545
546impl OnnxModel {
547    /// Create a new empty ONNX model.
548    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    /// Encode to ONNX protobuf binary format.
561    pub fn to_bytes(&self) -> Vec<u8> {
562        // Build GraphProto
563        let mut graph = PbEncoder::new();
564
565        // field 1: nodes (repeated NodeProto)
566        for node in &self.nodes {
567            let node_bytes = node.encode();
568            graph.write_bytes_field(1, &node_bytes);
569        }
570
571        // field 2: name
572        graph.write_string_field(2, &self.graph_name);
573
574        // field 5: initializers (repeated TensorProto — the weights)
575        for init in &self.initializers {
576            let tensor_bytes = init.encode();
577            graph.write_bytes_field(5, &tensor_bytes);
578        }
579
580        // field 11: inputs (repeated ValueInfoProto)
581        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        // field 12: outputs (repeated ValueInfoProto)
587        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        // Build ModelProto
593        let mut model = PbEncoder::new();
594        // field 1: ir_version (int64)
595        model.write_varint_field(1, ONNX_IR_VERSION as u64);
596        // field 2: producer_name
597        model.write_string_field(2, &self.producer_name);
598        // field 3: producer_version
599        model.write_string_field(3, &self.producer_version);
600        // field 7: graph (GraphProto)
601        model.write_message_field(7, &graph);
602        // field 8: opset_import (OperatorSetIdProto)
603        let mut opset = PbEncoder::new();
604        opset.write_string_field(1, ONNX_DOMAIN); // domain
605        opset.write_varint_field(2, ONNX_OPSET_VERSION as u64); // version
606        model.write_message_field(8, &opset);
607
608        model.into_bytes()
609    }
610
611    /// Save ONNX model to a file.
612    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
619/// Encode a ValueInfoProto (input/output description).
620fn encode_value_info(name: &str, dims: &[i64], data_type: i32) -> PbEncoder {
621    let mut vi = PbEncoder::new();
622    vi.write_string_field(1, name); // field 1: name
623
624    // field 2: type (TypeProto)
625    let mut type_proto = PbEncoder::new();
626    // field 1: tensor_type (Tensor_TypeProto)
627    let mut tensor_type = PbEncoder::new();
628    tensor_type.write_varint_field(1, data_type as u64); // elem_type
629                                                         // field 2: shape (TensorShapeProto)
630    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); // dim_value
635        } else {
636            dim.write_string_field(2, "dynamic"); // dim_param (symbolic)
637        }
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
647// =============================================================================
648// Export API
649// =============================================================================
650
651/// Export a module's weights as an ONNX model file.
652///
653/// This creates a "weight-only" ONNX model: the initializer tensors contain
654/// the model's learned parameters, and the graph describes a simple
655/// sequential pass from input to output.
656///
657/// # Arguments
658/// - `path`: output file path (typically `.onnx`)
659/// - `module`: the trained module to export
660/// - `model_name`: name for the ONNX graph
661/// - `input_shape`: shape of the model's input tensor
662///
663/// # Example
664/// ```ignore
665/// let model = Linear::new(784, 10, true, DType::F32, &dev)?;
666/// export_weights("model.onnx", &model, "classifier", &[1, 784])?;
667/// ```
668pub 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    // Add input
684    model
685        .inputs
686        .push(("input".to_string(), input_shape.to_vec(), ONNX_FLOAT));
687
688    // Add each parameter as an initializer
689    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                // Store as F32 for compatibility
706                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    // Add a simple identity graph: input → output through the params
715    // (The actual computation graph would require op-level tracking)
716    model.outputs.push((
717        "output".to_string(),
718        vec![-1], // dynamic output shape
719        ONNX_FLOAT,
720    ));
721
722    model.save(path)
723}
724
725/// Export named tensors directly to ONNX format.
726///
727/// Lower-level API: saves a set of named tensors as ONNX initializers.
728pub 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
766/// Convert a `shrew_ir::graph::IrGraph` into an `OnnxModel`.
767///
768/// Maps graph inputs, graph outputs, and internal computation nodes into
769/// standard ONNX nodes (e.g. `MatMul`, `Add`, `Relu`, `Softmax`, etc.).
770pub fn ir_graph_to_onnx(graph: &IrGraph) -> Result<OnnxModel> {
771    let mut model = OnnxModel::new(&graph.name);
772
773    // Inputs
774    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    // Outputs
804    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    // Nodes
834    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
935/// Export an entire `shrew_ir::graph::IrGraph` as an ONNX model file.
936pub 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
941// =============================================================================
942// Import API
943// =============================================================================
944
945/// Load tensor weights from an ONNX model file.
946///
947/// Returns a map of tensor name → Tensor for all initializers found.
948///
949/// # Example
950/// ```ignore
951/// let weights = load_onnx_weights::<CpuBackend>("model.onnx", &CpuDevice)?;
952/// for (name, tensor) in &weights {
953///     println!("{}: {:?}", name, tensor.dims());
954/// }
955/// ```
956pub 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
966/// Load tensor weights from ONNX bytes (in-memory).
967pub 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    // Parse ModelProto
975    while dec.remaining() > 0 {
976        let (field, wire_type) = dec.read_tag()?;
977        match (field, wire_type) {
978            (7, 2) => {
979                // GraphProto
980                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
993/// Parse graph initializer tensors from a GraphProto.
994fn 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                // Initializer (TensorProto)
1006                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// =============================================================================
1030// Graph Import — Full ONNX graph parsing
1031// =============================================================================
1032
1033/// A fully parsed ONNX graph: nodes + initializers + I/O metadata.
1034#[derive(Debug, Clone)]
1035pub struct OnnxGraph {
1036    /// Computation nodes in topological order.
1037    pub nodes: Vec<OnnxNode>,
1038    /// Initializer tensors (weights / constants).
1039    pub initializer_protos: Vec<OnnxTensor>,
1040    /// Graph input names (including initializer names).
1041    pub input_names: Vec<String>,
1042    /// Graph output names.
1043    pub output_names: Vec<String>,
1044    /// Graph name.
1045    pub name: String,
1046}
1047
1048/// Decode an OnnxNode from protobuf bytes.
1049fn 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
1076/// Decode an OnnxAttribute from protobuf bytes.
1077fn 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()?,           // name
1090            (2, 0) => attr_type = dec.read_varint()?,      // type
1091            (3, 0) => int_val = dec.read_varint()? as i64, // i
1092            (4, 5) => {
1093                // f (fixed32)
1094                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(), // s
1102            (7, 5) => {
1103                // floats (repeated fixed32)
1104                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                // floats (packed)
1113                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), // ints (repeated varint)
1119            (8, 2) => {
1120                // ints (packed)
1121                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), // fallback
1137    };
1138    Ok((name, val))
1139}
1140
1141/// Parse a full GraphProto: nodes, initializers, inputs, outputs.
1142fn 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                // input ValueInfoProto — extract name (field 1)
1167                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                // output ValueInfoProto
1173                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
1183/// Extract just the name from a ValueInfoProto.
1184fn 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
1196/// Load a full ONNX graph (nodes + initializers) from a file.
1197pub 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
1203/// Load a full ONNX graph from in-memory bytes.
1204pub 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
1217// =============================================================================
1218// Graph Execution — Run an ONNX graph with Shrew tensors
1219// =============================================================================
1220
1221/// Execute an ONNX graph on the given backend.
1222///
1223/// Takes a parsed `OnnxGraph` and a map of input tensors. Initializer tensors
1224/// from the graph are materialised on the given device. Each node is executed
1225/// in order (the ONNX spec requires nodes in topological order).
1226///
1227/// Returns a map of output-name → Tensor for all graph outputs.
1228///
1229/// # Supported ops
1230///
1231/// `Add`, `Sub`, `Mul`, `Div`, `MatMul`, `Gemm`, `Relu`, `Sigmoid`, `Tanh`,
1232/// `Softmax`, `LogSoftmax`, `Reshape`, `Transpose`, `Flatten`, `Squeeze`,
1233/// `Unsqueeze`, `Concat`, `Identity`, `Neg`, `Sqrt`, `Exp`, `Log`, `Abs`,
1234/// `Clip`, `ReduceMean`, `ReduceSum`, `ReduceMax`, `ReduceMin`, `Gather`,
1235/// `BatchNormalization`, `Dropout`, `Shape`, `Cast`, `Pow`.
1236///
1237/// Unsupported ops produce an error.
1238pub 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    // 1. Load initializers
1246    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    // 2. Insert user-provided inputs (overrides initializers if names clash)
1260    for (name, tensor) in inputs {
1261        env.insert(name.clone(), tensor.clone());
1262    }
1263
1264    // 3. Execute nodes in order
1265    for node in &graph.nodes {
1266        execute_node(node, &mut env, device)?;
1267    }
1268
1269    // 4. Collect outputs
1270    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
1279/// Helper: get a tensor from the environment by name.
1280fn 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
1288/// Helper: get an integer attribute with a default.
1289fn 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
1296/// Helper: get an integer-list attribute.
1297fn 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
1304/// Helper: get a float attribute with default.
1305fn 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
1312/// Execute a single ONNX node, inserting results into the environment.
1313fn 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        // ── Element-wise binary ──────────────────────────────────────────
1320        "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            // ONNX Pow has two inputs; exponent is second input
1347            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 / Gemm ────────────────────────────────────────────────
1360        "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            // Y = alpha * A' * B' + beta * C
1368            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        // ── Unary activations ────────────────────────────────────────────
1400        "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 / LogSoftmax ─────────────────────────────────────────
1434        "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 (clamp) ─────────────────────────────────────────────────
1456        "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        // ── Shape manipulation ───────────────────────────────────────────
1472        "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            // Resolve -1 dims
1478            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); // -1 as usize wraps
1481            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                // Default: reverse all dims
1499                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                // Squeeze from highest axis to lowest to avoid index shifting
1520                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                // ONNX opset >= 13: axes is a tensor input
1533                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        // ── Reductions ───────────────────────────────────────────────────
1572        "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 ───────────────────────────────────────────────────────
1634        "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 ───────────────────────────────────────────
1642        "BatchNormalization" => {
1643            // inputs: X, scale, B, mean, var
1644            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            // y = scale * (x - mean) / sqrt(var + eps) + bias
1652            // Broadcast: mean/var/scale/bias are 1-D [C], x is [N, C, ...]
1653            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 (inference) ──────────────────────────────────────────
1662        "Dropout" => {
1663            // In inference mode, Dropout is identity
1664            let x = get_tensor(env, &node.inputs[0])?.clone();
1665            env.insert(node.outputs[0].clone(), x.clone());
1666            // Optional second output (mask) — insert copy
1667            if node.outputs.len() > 1 && !node.outputs[1].is_empty() {
1668                env.insert(node.outputs[1].clone(), x);
1669            }
1670        }
1671
1672        // ── Identity ─────────────────────────────────────────────────────
1673        "Identity" => {
1674            let x = get_tensor(env, &node.inputs[0])?;
1675            env.insert(node.outputs[0].clone(), x.clone());
1676        }
1677
1678        // ── Shape ────────────────────────────────────────────────────────
1679        "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 ─────────────────────────────────────────────────────────
1688        "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 ─────────────────────────────────────────────────────
1696        "Constant" => {
1697            // Try to get value from attributes
1698            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// =============================================================================
1733// Tests
1734// =============================================================================
1735
1736#[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        // Data stored in raw_data (packed float), so check via to_f64_vec
1778        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        // Export
1790        let path = std::env::temp_dir().join("shrew_test_onnx.onnx");
1791        export_weights(&path, &linear, "test_model", &[1, 4]).unwrap();
1792
1793        // Import
1794        let weights = load_onnx_weights::<B>(&path, &DEV).unwrap();
1795
1796        // Should have weight and bias
1797        assert_eq!(weights.len(), 2);
1798
1799        // Verify weight shape
1800        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            // Values should match
1806            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        // Cleanup
1814        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); // non-trivial size
1869    }
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    // ────────────────────────────────────────────────────────────────────
1903    //  Graph Import + Execution tests
1904    // ────────────────────────────────────────────────────────────────────
1905
1906    /// Helper: build a minimal ONNX model bytes from an OnnxModel.
1907    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        // Graph:  Y = A + B
1915        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        // Graph:  Z = Relu(X * W + B)
1948        //   matmul_out = MatMul(X, W)
1949        //   add_out    = Add(matmul_out, B)
1950        //   Z          = Relu(add_out)
1951        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        // W and B as initializers
1956        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        // X = [[1.0, -1.0]]
1994        // X*W = [[1*1+(-1)*0, 1*(-1)+(-1)*2, 1*0.5+(-1)*(-0.5)]]
1995        //     = [[1.0, -3.0, 1.0]]
1996        // Relu => [[1.0, 0.0, 1.0]]
1997        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        // Build, save, load from file, execute
2044        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); // sigmoid(0) = 0.5
2069        assert!((data[1] - 1.0).abs() < 1e-3); // sigmoid(1000) ≈ 1.0
2070
2071        let _ = fs::remove_file(&path);
2072    }
2073
2074    #[test]
2075    fn test_decode_attribute_roundtrip() {
2076        // Encode an Int attribute, decode it back
2077        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}