1use crate::error::{PyResult, RonnError};
4use crate::session::PySession;
5use pyo3::prelude::*;
6use ronn_api::Model;
7
8#[pyclass(name = "Model")]
18pub struct PyModel {
19 inner: Model,
20}
21
22#[pymethods]
23impl PyModel {
24 #[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 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 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 #[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 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 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 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 let session = self
132 .inner
133 .create_session(options)
134 .map_err(RonnError::from)?;
135
136 Ok(PySession::new(session))
137 }
138
139 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}