Skip to main content

runmat_runtime/builtins/common/
shape.rs

1use runmat_builtins::{Tensor, Value};
2
3use crate::dispatcher::gather_if_needed_async;
4use crate::RuntimeError;
5
6/// Return true if a shape should be treated as a scalar.
7pub fn is_scalar_shape(shape: &[usize]) -> bool {
8    shape.is_empty()
9        || (shape.len() == 1 && shape[0] == 1)
10        || (shape.len() == 2 && shape[0] == 1 && shape[1] == 1)
11}
12
13/// Return the canonical scalar shape.
14pub fn canonical_scalar_shape() -> Vec<usize> {
15    vec![1, 1]
16}
17
18/// Normalize scalar-like shapes to the canonical scalar shape.
19pub fn normalize_scalar_shape(shape: &[usize]) -> Vec<usize> {
20    if is_scalar_shape(shape) {
21        canonical_scalar_shape()
22    } else {
23        shape.to_vec()
24    }
25}
26
27/// Normalize a raw shape vector into MATLAB-compatible dimension metadata.
28fn normalize_shape(shape: &[usize]) -> Vec<usize> {
29    if shape.len() == 1 && shape[0] != 1 {
30        return vec![1, shape[0]];
31    }
32    if is_scalar_shape(shape) {
33        return canonical_scalar_shape();
34    }
35    shape.to_vec()
36}
37
38/// Return the MATLAB-visible dimension vector for a runtime value.
39#[async_recursion::async_recursion(?Send)]
40pub async fn value_dimensions(value: &Value) -> Result<Vec<usize>, RuntimeError> {
41    let dims = match value {
42        Value::Tensor(t) => normalize_shape(&t.shape),
43        Value::SparseTensor(t) => normalize_shape(&[t.rows, t.cols]),
44        Value::ComplexTensor(t) => normalize_shape(&t.shape),
45        Value::LogicalArray(la) => normalize_shape(&la.shape),
46        Value::StringArray(sa) => normalize_shape(&sa.shape),
47        Value::SymbolicArray(sa) => normalize_shape(&sa.shape),
48        Value::CharArray(ca) => vec![ca.rows, ca.cols],
49        Value::Cell(ca) => normalize_shape(&ca.shape),
50        Value::GpuTensor(handle) => {
51            if handle.shape.is_empty() {
52                let gathered = gather_if_needed_async(&Value::GpuTensor(handle.clone())).await?;
53                return value_dimensions(&gathered).await;
54            }
55            normalize_shape(&handle.shape)
56        }
57        _ => vec![1, 1],
58    };
59    Ok(dims)
60}
61
62/// Compute the total number of elements contained in a runtime value.
63#[async_recursion::async_recursion(?Send)]
64pub async fn value_numel(value: &Value) -> Result<usize, RuntimeError> {
65    let numel = match value {
66        Value::Tensor(t) => t.data.len(),
67        Value::SparseTensor(t) => t.rows.saturating_mul(t.cols),
68        Value::ComplexTensor(t) => t.data.len(),
69        Value::LogicalArray(la) => la.data.len(),
70        Value::StringArray(sa) => sa.data.len(),
71        Value::SymbolicArray(sa) => sa.data.len(),
72        Value::CharArray(ca) => ca.rows * ca.cols,
73        Value::Cell(ca) => ca.data.len(),
74        Value::GpuTensor(handle) => {
75            if handle.shape.is_empty() {
76                let gathered = gather_if_needed_async(&Value::GpuTensor(handle.clone())).await?;
77                return value_numel(&gathered).await;
78            }
79            handle
80                .shape
81                .iter()
82                .copied()
83                .fold(1usize, |acc, dim| acc.saturating_mul(dim))
84        }
85        _ => 1,
86    };
87    Ok(numel)
88}
89
90/// Compute the dimensionality (NDIMS) of a runtime value, with MATLAB semantics.
91pub async fn value_ndims(value: &Value) -> Result<usize, RuntimeError> {
92    let dims = value_dimensions(value).await?;
93    if dims.len() < 2 {
94        Ok(2)
95    } else {
96        Ok(dims.len())
97    }
98}
99
100/// Convert a dimension vector into a 1×N tensor encoded as `f64`.
101pub fn dims_to_row_tensor(dims: &[usize]) -> Result<Tensor, String> {
102    let len = dims.len();
103    let data: Vec<f64> = dims.iter().map(|&d| d as f64).collect();
104    let shape = if len == 0 { vec![1, 0] } else { vec![1, len] };
105    Tensor::new(data, shape).map_err(|e| format!("shape::dims_to_row_tensor: {e}"))
106}
107
108#[cfg(test)]
109pub(crate) mod tests {
110    use super::*;
111    use futures::executor::block_on;
112
113    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
114    #[test]
115    fn dims_scalar_defaults_to_one_by_one() {
116        assert_eq!(
117            block_on(value_dimensions(&Value::Num(5.0))).unwrap(),
118            vec![1, 1]
119        );
120    }
121
122    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
123    #[test]
124    fn dims_tensor_preserves_rank() {
125        let tensor = Tensor::new(vec![0.0; 12], vec![2, 3, 2]).unwrap();
126        assert_eq!(
127            block_on(value_dimensions(&Value::Tensor(tensor))).unwrap(),
128            vec![2, 3, 2]
129        );
130    }
131
132    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
133    #[test]
134    fn numel_gpu_uses_shape_product() {
135        let handle = runmat_accelerate_api::GpuTensorHandle {
136            shape: vec![4, 5, 6],
137            device_id: 0,
138            buffer_id: 1,
139        };
140        assert_eq!(
141            block_on(value_numel(&Value::GpuTensor(handle))).unwrap(),
142            120
143        );
144    }
145
146    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
147    #[test]
148    fn dims_to_row_tensor_converts() {
149        let tensor = dims_to_row_tensor(&[2, 4, 6]).unwrap();
150        assert_eq!(tensor.shape, vec![1, 3]);
151        assert_eq!(tensor.data, vec![2.0, 4.0, 6.0]);
152    }
153}