Skip to main content

ronn/
model.rs

1//! Model class for Python bindings
2
3use crate::error::{PyResult, RonnError};
4use crate::session::PySession;
5use pyo3::prelude::*;
6use ronn_api::Model;
7
8/// ONNX model
9///
10/// # Example
11///
12/// ```python
13/// model = ronn.Model.load("model.onnx")
14/// print(f"Inputs: {model.input_names()}")
15/// print(f"Outputs: {model.output_names()}")
16/// ```
17#[pyclass(name = "Model")]
18pub struct PyModel {
19    inner: Model,
20}
21
22#[pymethods]
23impl PyModel {
24    /// Load a model from an ONNX file
25    ///
26    /// # Arguments
27    ///
28    /// * `path` - Path to .onnx file
29    ///
30    /// # Example
31    ///
32    /// ```python
33    /// model = ronn.Model.load("resnet18.onnx")
34    /// ```
35    #[staticmethod]
36    fn load(path: &str) -> PyResult<Self> {
37        let model = Model::load(path).map_err(RonnError::from)?;
38        Ok(Self { inner: model })
39    }
40
41    /// Get model input names
42    ///
43    /// # Returns
44    ///
45    /// List of input tensor names
46    fn input_names(&self) -> Vec<String> {
47        self.inner
48            .input_names()
49            .iter()
50            .map(|s| s.to_string())
51            .collect()
52    }
53
54    /// Get model output names
55    ///
56    /// # Returns
57    ///
58    /// List of output tensor names
59    fn output_names(&self) -> Vec<String> {
60        self.inner
61            .output_names()
62            .iter()
63            .map(|s| s.to_string())
64            .collect()
65    }
66
67    /// Create an inference session
68    ///
69    /// # Arguments
70    ///
71    /// * `optimization_level` - Graph optimization level ("none", "basic", "aggressive")
72    /// * `provider` - Execution provider ("cpu", "gpu", "bitnet", "wasm")
73    /// * `num_threads` - Number of threads for CPU execution
74    ///
75    /// # Example
76    ///
77    /// ```python
78    /// session = model.create_session(
79    ///     optimization_level="aggressive",
80    ///     provider="cpu",
81    ///     num_threads=4
82    /// )
83    /// ```
84    #[pyo3(signature = (optimization_level="basic", provider="cpu", num_threads=None))]
85    fn create_session(
86        &self,
87        optimization_level: &str,
88        provider: &str,
89        num_threads: Option<usize>,
90    ) -> PyResult<PySession> {
91        use ronn_api::SessionOptions;
92
93        // Parse optimization level
94        let opt_level = match optimization_level {
95            "O0" | "none" => ronn_graph::OptimizationLevel::O0,
96            "O1" | "basic" => ronn_graph::OptimizationLevel::O1,
97            "O2" | "default" => ronn_graph::OptimizationLevel::O2,
98            "O3" | "aggressive" => ronn_graph::OptimizationLevel::O3,
99            _ => {
100                return Err(RonnError(format!(
101                    "Invalid optimization level: {}. Use 'none'/'O0', 'basic'/'O1', 'default'/'O2', or 'aggressive'/'O3'",
102                    optimization_level
103                )));
104            }
105        };
106
107        // Parse provider
108        let provider_type = match provider {
109            "cpu" => ronn_core::ProviderId::CPU,
110            "gpu" | "cuda" => ronn_core::ProviderId::GPU,
111            "bitnet" => ronn_core::ProviderId::BitNet,
112            "wasm" => ronn_core::ProviderId::WebAssembly,
113            _ => {
114                return Err(RonnError(format!(
115                    "Invalid provider: {}. Use 'cpu', 'gpu', 'bitnet', or 'wasm'",
116                    provider
117                )));
118            }
119        };
120
121        // Build session options
122        let mut options = SessionOptions::default()
123            .with_optimization_level(opt_level)
124            .with_provider(provider_type);
125
126        if let Some(threads) = num_threads {
127            options = options.with_num_threads(threads);
128        }
129
130        // Create session
131        let session = self
132            .inner
133            .create_session(options)
134            .map_err(RonnError::from)?;
135
136        Ok(PySession::new(session))
137    }
138
139    /// Get model metadata
140    ///
141    /// # Returns
142    ///
143    /// Dictionary with model information
144    fn metadata(&self, py: Python) -> PyObject {
145        let dict = pyo3::types::PyDict::new(py);
146        dict.set_item("inputs", self.input_names()).unwrap();
147        dict.set_item("outputs", self.output_names()).unwrap();
148        dict.into()
149    }
150
151    fn __repr__(&self) -> String {
152        format!(
153            "Model(inputs={:?}, outputs={:?})",
154            self.input_names(),
155            self.output_names()
156        )
157    }
158}