sc_neurocore_engine/bindings/
matrix_inputs.rs1use pyo3::exceptions::PyValueError;
12use pyo3::prelude::*;
13
14pub(crate) fn extract_matrix_f64(
15 data: &Bound<'_, PyAny>,
16 name: &str,
17) -> PyResult<(Vec<f64>, usize, usize)> {
18 if let Ok(rows) = data.extract::<Vec<Vec<f64>>>() {
19 if rows.is_empty() {
20 return Err(PyValueError::new_err(format!(
21 "{} must not be an empty matrix.",
22 name
23 )));
24 }
25 let row_count = rows.len();
26 let cols = rows[0].len();
27 if cols == 0 {
28 return Err(PyValueError::new_err(format!(
29 "{} must not have zero columns.",
30 name
31 )));
32 }
33 if rows.iter().any(|r| r.len() != cols) {
34 return Err(PyValueError::new_err(format!(
35 "{} must be a rectangular matrix.",
36 name
37 )));
38 }
39 let out = rows.into_iter().flatten().collect::<Vec<f64>>();
40 return Ok((out, row_count, cols));
41 }
42
43 if let Ok(flat) = data.extract::<Vec<f64>>() {
44 if flat.is_empty() {
45 return Err(PyValueError::new_err(format!(
46 "{} must not be an empty vector.",
47 name
48 )));
49 }
50 let cols = flat.len();
51 return Ok((flat, 1, cols));
52 }
53
54 Err(PyValueError::new_err(format!(
55 "{} must be a 1-D or 2-D float array.",
56 name
57 )))
58}
59
60pub(crate) fn reshape_flat_to_rows(flat: Vec<f64>, rows: usize, cols: usize) -> Vec<Vec<f64>> {
61 let mut out = Vec::with_capacity(rows);
62 for i in 0..rows {
63 out.push(flat[i * cols..(i + 1) * cols].to_vec());
64 }
65 out
66}