Skip to main content

runmat_runtime/builtins/common/
tensor.rs

1use std::{borrow::Cow, convert::TryFrom};
2
3use num_complex::Complex64;
4use runmat_accelerate_api::{GpuTensorStorage, HostTensorOwned};
5use runmat_value::{
6    ComplexTensor, IntValue, IntegerStorage, LogicalArray, NumericDType, NumericStorage, Tensor,
7    Value,
8};
9
10use crate::dispatcher::gather_if_needed_async;
11
12/// Return the total number of elements for a given shape.
13pub fn element_count(shape: &[usize]) -> usize {
14    let mut acc: u128 = 1;
15    for &dim in shape {
16        let dim128 = dim as u128;
17        acc = acc
18            .checked_mul(dim128)
19            .expect("tensor::element_count: overflow computing element count");
20    }
21    usize::try_from(acc).expect("tensor::element_count: overflow converting to usize")
22}
23
24/// Construct a zero-filled tensor with the provided shape.
25pub fn zeros(shape: &[usize]) -> Result<Tensor, String> {
26    Tensor::new(vec![0.0; element_count(shape)], shape.to_vec())
27        .map_err(|e| format!("tensor zeros: {e}"))
28}
29
30/// Construct an one-filled tensor with the provided shape.
31pub fn ones(shape: &[usize]) -> Result<Tensor, String> {
32    Tensor::new(vec![1.0; element_count(shape)], shape.to_vec())
33        .map_err(|e| format!("tensor ones: {e}"))
34}
35
36/// Construct a zero-filled tensor with an explicit dtype flag.
37pub fn zeros_with_dtype(shape: &[usize], dtype: NumericDType) -> Result<Tensor, String> {
38    integer_tensor_with_value(shape, dtype, false)
39        .unwrap_or_else(|| {
40            Tensor::new_with_dtype(vec![0.0; element_count(shape)], shape.to_vec(), dtype)
41        })
42        .map_err(|e| format!("tensor zeros: {e}"))
43}
44
45/// Construct a one-filled tensor with an explicit dtype flag.
46pub fn ones_with_dtype(shape: &[usize], dtype: NumericDType) -> Result<Tensor, String> {
47    integer_tensor_with_value(shape, dtype, true)
48        .unwrap_or_else(|| {
49            Tensor::new_with_dtype(vec![1.0; element_count(shape)], shape.to_vec(), dtype)
50        })
51        .map_err(|e| format!("tensor ones: {e}"))
52}
53
54fn integer_tensor_with_value(
55    shape: &[usize],
56    dtype: NumericDType,
57    ones: bool,
58) -> Option<Result<Tensor, String>> {
59    let len = element_count(shape);
60    let storage = match dtype {
61        NumericDType::I8 => IntegerStorage::I8(vec![if ones { 1 } else { 0 }; len]),
62        NumericDType::I16 => IntegerStorage::I16(vec![if ones { 1 } else { 0 }; len]),
63        NumericDType::I32 => IntegerStorage::I32(vec![if ones { 1 } else { 0 }; len]),
64        NumericDType::I64 => IntegerStorage::I64(vec![if ones { 1 } else { 0 }; len]),
65        NumericDType::U8 => IntegerStorage::U8(vec![if ones { 1 } else { 0 }; len]),
66        NumericDType::U16 => IntegerStorage::U16(vec![if ones { 1 } else { 0 }; len]),
67        NumericDType::U32 => IntegerStorage::U32(vec![if ones { 1 } else { 0 }; len]),
68        NumericDType::U64 => IntegerStorage::U64(vec![if ones { 1 } else { 0 }; len]),
69        NumericDType::F32 | NumericDType::F64 => return None,
70    };
71    Some(Tensor::new_integer(storage, shape.to_vec()))
72}
73
74/// Converts floating-point values to an exact integer tensor using the
75/// prototype class's MATLAB assignment semantics.
76pub fn integer_tensor_from_f64_like(
77    prototype: &IntegerStorage,
78    values: Vec<f64>,
79    shape: &[usize],
80) -> Result<Tensor, String> {
81    let storage = prototype
82        .from_same_class_values(
83            values
84                .into_iter()
85                .map(|value| prototype.cast_f64_assignment(value))
86                .collect(),
87        )
88        .map_err(|e| format!("integer tensor conversion: {e}"))?;
89    Tensor::new_integer(storage, shape.to_vec())
90        .map_err(|e| format!("integer tensor conversion: {e}"))
91}
92
93/// Materialize a tensor's authoritative numeric values in the f64 computation domain.
94/// Callers that admit integers are responsible for any required exactness check.
95pub fn tensor_values_f64(tensor: &Tensor) -> Vec<f64> {
96    tensor.materialize_f64()
97}
98
99/// Return a borrowed view for native double tensors and explicitly materialize
100/// every other numeric class in the f64 computation domain.
101pub fn tensor_values_f64_cow(tensor: &Tensor) -> Cow<'_, [f64]> {
102    match tensor.as_f64_slice() {
103        Some(values) => Cow::Borrowed(values),
104        None => Cow::Owned(tensor.materialize_f64()),
105    }
106}
107
108/// Read one authoritative numeric tensor value into the f64 computation domain.
109/// Callers that admit integers are responsible for any required exactness check.
110pub fn tensor_value_f64(tensor: &Tensor, index: usize) -> f64 {
111    match tensor
112        .numeric_value_at(index)
113        .expect("tensor_value_f64: numeric storage index is in bounds")
114    {
115        runmat_value::NumericScalar::F64(value) => value,
116        runmat_value::NumericScalar::F32(value) => f64::from(value),
117        scalar => scalar
118            .into_int_value()
119            .expect("non-floating numeric scalar is integer")
120            .to_f64(),
121    }
122}
123
124/// Return an exact integer scalar from a scalar integer value or scalar typed
125/// integer tensor.
126pub fn scalar_integer_value(value: &Value) -> Option<IntValue> {
127    match value {
128        Value::Int(value) => Some(value.clone()),
129        Value::Tensor(tensor) if is_scalar_tensor(tensor) => tensor
130            .integer_storage()
131            .and_then(|storage| storage.value_at(0)),
132        _ => None,
133    }
134}
135
136/// Consume a tensor and explicitly materialize its authoritative storage in
137/// the f64 computation domain.
138pub fn tensor_into_values_f64(tensor: Tensor) -> Vec<f64> {
139    tensor
140        .into_numeric_storage()
141        .expect("validated tensor storage")
142        .materialize_f64()
143}
144
145/// Consume a tensor at a provider API boundary that currently accepts only owned f64 storage.
146pub fn tensor_into_host_f64_owned(tensor: Tensor) -> HostTensorOwned {
147    let shape = tensor.shape.clone();
148    HostTensorOwned {
149        data: tensor_into_values_f64(tensor),
150        shape,
151        storage: GpuTensorStorage::Real,
152    }
153}
154
155/// Return a complex tensor's numeric values as Complex64, reading typed integer
156/// real/imaginary storage exactly instead of using the compatibility buffer.
157pub fn complex_tensor_values_complex64(tensor: &ComplexTensor) -> Vec<Complex64> {
158    tensor
159        .integer_storage()
160        .as_ref()
161        .map(|storage| {
162            let real = storage.real.exact_values();
163            let imag = storage.imag.exact_values();
164            real.into_iter()
165                .zip(imag)
166                .map(|(re, im)| Complex64::new(re.to_f64(), im.to_f64()))
167                .collect()
168        })
169        .unwrap_or_else(|| {
170            tensor
171                .materialize_f64()
172                .iter()
173                .map(|&(re, im)| Complex64::new(re, im))
174                .collect()
175        })
176}
177
178/// Return a complex tensor's element count, reading exact integer-complex
179/// storage instead of the floating compatibility buffer when present.
180pub fn complex_tensor_element_len(tensor: &ComplexTensor) -> usize {
181    tensor
182        .integer_storage()
183        .as_ref()
184        .map_or(tensor.materialize_f64().len(), |storage| storage.len())
185}
186
187/// Return true when a complex tensor contains exactly one scalar element.
188pub fn is_scalar_complex_tensor(tensor: &ComplexTensor) -> bool {
189    complex_tensor_element_len(tensor) == 1
190}
191
192/// Return one complex tensor value as Complex64, reading typed integer-complex
193/// storage exactly instead of using the compatibility backing buffer.
194pub fn complex_tensor_value_complex64(tensor: &ComplexTensor, index: usize) -> Complex64 {
195    match tensor.integer_storage() {
196        Some(storage) => {
197            let real = storage
198                .real
199                .value_at(index)
200                .expect("complex_tensor_value_complex64: real storage index is in bounds")
201                .to_f64();
202            let imag = storage
203                .imag
204                .value_at(index)
205                .expect("complex_tensor_value_complex64: imaginary storage index is in bounds")
206                .to_f64();
207            Complex64::new(real, imag)
208        }
209        None => {
210            let (real, imag) = tensor.materialize_f64()[index];
211            Complex64::new(real, imag)
212        }
213    }
214}
215
216/// Consume a complex tensor and return Complex64 values, preserving the fast
217/// path for ordinary complex double tensors.
218pub fn complex_tensor_into_values_complex64(tensor: ComplexTensor) -> Vec<Complex64> {
219    if tensor.integer_storage().is_some() {
220        complex_tensor_values_complex64(&tensor)
221    } else {
222        tensor
223            .materialize_f64()
224            .into_iter()
225            .map(|(re, im)| Complex64::new(re, im))
226            .collect()
227    }
228}
229
230/// Normalize typed integer tensors to f64 at builtin boundaries that
231/// intentionally compute double-valued numeric/statistical results.
232pub fn integer_tensor_to_f64(tensor: Tensor) -> Result<Tensor, String> {
233    if tensor.integer_storage().is_none() {
234        return Ok(tensor);
235    }
236    Tensor::new(tensor_values_f64(&tensor), tensor.shape.clone())
237        .map_err(|e| format!("integer tensor conversion: {e}"))
238}
239
240/// Convert a logical array (0/1 bytes) into a numeric tensor.
241pub fn logical_to_tensor(logical: &LogicalArray) -> Result<Tensor, String> {
242    let data: Vec<f64> = logical
243        .data
244        .iter()
245        .map(|&b| if b != 0 { 1.0 } else { 0.0 })
246        .collect();
247    Tensor::new(data, logical.shape.clone()).map_err(|e| format!("logical->tensor: {e}"))
248}
249
250fn value_into_tensor_impl(name: &str, value: Value) -> Result<Tensor, String> {
251    match value {
252        Value::Tensor(t) => Ok(t),
253        Value::LogicalArray(logical) => logical_to_tensor(&logical),
254        Value::Num(n) => Tensor::new(vec![n], vec![1, 1]).map_err(|e| format!("tensor: {e}")),
255        Value::Int(i) => Tensor::new_integer(IntegerStorage::from_scalar(i), vec![1, 1])
256            .map_err(|e| format!("tensor: {e}")),
257        Value::Bool(b) => Tensor::new(vec![if b { 1.0 } else { 0.0 }], vec![1, 1])
258            .map_err(|e| format!("tensor: {e}")),
259        other => Err(format!(
260            "{name}: unsupported input type {:?}; expected numeric or logical values",
261            other
262        )),
263    }
264}
265
266/// Convert a `Value` into an owned `Tensor`, defaulting error messages to `"sum"`.
267pub fn value_into_tensor(value: Value) -> Result<Tensor, String> {
268    value_into_tensor_impl("sum", value)
269}
270
271/// Convert a `Value` into a tensor while customising the builtin name in error messages.
272pub fn value_into_tensor_for(name: &str, value: Value) -> Result<Tensor, String> {
273    value_into_tensor_impl(name, value)
274}
275
276/// Clone a `Value` and coerce it into a tensor.
277pub fn value_to_tensor(value: &Value) -> Result<Tensor, String> {
278    value_into_tensor(value.clone())
279}
280
281/// Convert a `Tensor` back into a runtime value.
282///
283/// Scalars (exactly one element) become their exact scalar representation;
284/// all other tensors remain as dense tensor variants.
285pub fn tensor_into_value(tensor: Tensor) -> Value {
286    if is_scalar_tensor(&tensor) {
287        if let Some(storage) = tensor.integer_storage() {
288            return Value::Int(storage.value_at(0).expect("one-element integer storage"));
289        }
290        if tensor.numeric_dtype() == runmat_value::NumericDType::F64 {
291            Value::Num(tensor_value_f64(&tensor, 0))
292        } else {
293            Value::Tensor(tensor)
294        }
295    } else {
296        Value::Tensor(tensor)
297    }
298}
299
300/// Return true when a tensor contains exactly one scalar element.
301pub fn is_scalar_tensor(tensor: &Tensor) -> bool {
302    tensor_element_len(tensor) == 1
303}
304
305pub fn tensor_element_len(tensor: &Tensor) -> usize {
306    tensor.len()
307}
308
309fn scalar_f64_from_host_value(value: &Value) -> Result<Option<f64>, String> {
310    match value {
311        Value::Num(n) => Ok(Some(*n)),
312        Value::Int(i) => Ok(Some(i.to_f64())),
313        Value::Bool(b) => Ok(Some(if *b { 1.0 } else { 0.0 })),
314        Value::Tensor(t) => {
315            if is_scalar_tensor(t) {
316                if let Some(storage) = t.integer_storage() {
317                    return Ok(Some(
318                        storage
319                            .value_at(0)
320                            .expect("one-element integer storage")
321                            .to_f64(),
322                    ));
323                }
324                Ok(Some(tensor_value_f64(t, 0)))
325            } else {
326                Err(format!(
327                    "expected scalar tensor, got tensor of size {}",
328                    tensor_element_len(t)
329                ))
330            }
331        }
332        Value::LogicalArray(la) => {
333            if la.data.len() == 1 {
334                Ok(Some(if la.data[0] != 0 { 1.0 } else { 0.0 }))
335            } else {
336                Err(format!(
337                    "expected scalar logical array, got array of size {}",
338                    la.data.len()
339                ))
340            }
341        }
342        _ => Ok(None),
343    }
344}
345
346/// Attempt to extract a scalar f64 from a runtime value asynchronously.
347pub async fn scalar_f64_from_value_async(value: &Value) -> Result<Option<f64>, String> {
348    match value {
349        Value::GpuTensor(handle) => {
350            if !handle.shape.is_empty() {
351                let len = element_count(&handle.shape);
352                if len != 1 {
353                    return Err(format!("expected scalar gpuArray, got array of size {len}"));
354                }
355            }
356            let gathered = gather_if_needed_async(&Value::GpuTensor(handle.clone()))
357                .await
358                .map_err(|e| format!("scalar: {e}"))?;
359            scalar_f64_from_host_value(&gathered)
360        }
361        _ => scalar_f64_from_host_value(value),
362    }
363}
364
365/// Attempt to parse a dimension index from a scalar-like runtime value.
366pub async fn dimension_from_value_async(
367    value: &Value,
368    name: &str,
369    allow_zero: bool,
370) -> Result<Option<usize>, String> {
371    match value {
372        Value::Int(value) => return parse_integer_dimension(value, name, allow_zero).map(Some),
373        Value::Tensor(tensor) if is_scalar_tensor(tensor) => {
374            if let Some(storage) = tensor.integer_storage() {
375                let value = storage.value_at(0).expect("one-element integer storage");
376                return parse_integer_dimension(&value, name, allow_zero).map(Some);
377            }
378        }
379        _ => {}
380    }
381    let Some(raw) = scalar_f64_from_value_async(value).await? else {
382        return Ok(None);
383    };
384    parse_numeric_dimension_value(raw, name, allow_zero).map(Some)
385}
386
387fn parse_integer_dimension(
388    value: &IntValue,
389    name: &str,
390    allow_zero: bool,
391) -> Result<usize, String> {
392    let dim = value
393        .try_to_usize()
394        .ok_or_else(|| format!("{name}: dimension is outside the supported range"))?;
395    if !allow_zero && dim == 0 {
396        return Err(format!("{name}: dimension must be >= 1"));
397    }
398    Ok(dim)
399}
400
401/// Parse every element of a typed integer tensor as a dimension selector without
402/// routing through the tensor's f64 compatibility backing store.
403pub fn integer_tensor_dimension_vector(
404    tensor: &Tensor,
405    name: &str,
406    allow_zero: bool,
407) -> Option<Result<Vec<usize>, String>> {
408    let storage = tensor.integer_storage()?;
409    Some(
410        (0..storage.len())
411            .map(|index| {
412                let value = storage
413                    .value_at(index)
414                    .expect("integer tensor storage length matches element count");
415                parse_integer_dimension(&value, name, allow_zero)
416            })
417            .collect(),
418    )
419}
420
421fn parse_integer_shape_dimension(value: &IntValue) -> Result<usize, String> {
422    value
423        .try_to_usize()
424        .ok_or_else(|| "dimensions must be non-negative platform integers".to_string())
425}
426
427fn parse_numeric_dimension(value: f64) -> Result<usize, String> {
428    if !value.is_finite() {
429        return Err("dimensions must be finite".to_string());
430    }
431    if value < 0.0 {
432        return Err("matrix dimensions must be non-negative".to_string());
433    }
434    let rounded = value.round();
435    if (rounded - value).abs() > f64::EPSILON {
436        return Err("dimensions must be integers".to_string());
437    }
438    if !fits_platform_usize(rounded) {
439        return Err("dimensions are outside the supported platform range".to_string());
440    }
441    Ok(rounded as usize)
442}
443
444fn fits_platform_usize(value: f64) -> bool {
445    value < usize::MAX as f64 || (usize::BITS < 64 && value == usize::MAX as f64)
446}
447
448fn dims_from_tensor_values(values: &[f64], shape: &[usize]) -> Result<Option<Vec<usize>>, String> {
449    let len = values.len();
450    if len == 0 {
451        return Ok(Some(Vec::new()));
452    }
453    let is_scalar = len == 1;
454    let is_row = shape.len() >= 2 && shape[0] == 1;
455    let is_column = shape.len() >= 2 && shape[1] == 1;
456    if !(is_row || is_column || is_scalar || shape.len() == 1) {
457        return Ok(None);
458    }
459    let mut dims = Vec::with_capacity(len);
460    for &value in values {
461        dims.push(parse_numeric_dimension(value)?);
462    }
463    Ok(Some(dims))
464}
465
466fn dims_from_integer_tensor_values(
467    storage: &IntegerStorage,
468    shape: &[usize],
469) -> Result<Option<Vec<usize>>, String> {
470    let len = storage.len();
471    if len == 0 {
472        return Ok(Some(Vec::new()));
473    }
474    let is_scalar = len == 1;
475    let is_row = shape.len() >= 2 && shape[0] == 1;
476    let is_column = shape.len() >= 2 && shape[1] == 1;
477    if !(is_row || is_column || is_scalar || shape.len() == 1) {
478        return Ok(None);
479    }
480    let mut dims = Vec::with_capacity(len);
481    for index in 0..len {
482        dims.push(parse_integer_shape_dimension(
483            &storage.value_at(index).expect("integer storage index"),
484        )?);
485    }
486    Ok(Some(dims))
487}
488
489/// Attempt to parse a dimension vector from a runtime value asynchronously.
490pub async fn dims_from_value_async(value: &Value) -> Result<Option<Vec<usize>>, String> {
491    match value {
492        Value::Num(n) => parse_numeric_dimension(*n).map(|dim| Some(vec![dim])),
493        Value::Int(i) => parse_integer_shape_dimension(i).map(|dim| Some(vec![dim])),
494        Value::Tensor(t) => match t.integer_storage() {
495            Some(storage) => dims_from_integer_tensor_values(storage, &t.shape),
496            None => dims_from_tensor_values(tensor_values_f64_cow(t).as_ref(), &t.shape),
497        },
498        Value::LogicalArray(la) => {
499            let values: Vec<f64> = la
500                .data
501                .iter()
502                .map(|&b| if b != 0 { 1.0 } else { 0.0 })
503                .collect();
504            dims_from_tensor_values(&values, &la.shape)
505        }
506        Value::GpuTensor(handle) => {
507            let gathered = gather_if_needed_async(&Value::GpuTensor(handle.clone()))
508                .await
509                .map_err(|e| format!("dimensions: {e}"))?;
510            match gathered {
511                Value::Tensor(t) => {
512                    if t.is_empty() {
513                        tracing::warn!(
514                            gpu_shape = ?handle.shape,
515                            "dims_from_value_async: gathered GPU tensor has no data"
516                        );
517                    }
518                    tracing::trace!(
519                        "dims_from_value_async: GPU tensor values gpu_shape={:?} host_shape={:?} class={} elements={}",
520                        handle.shape,
521                        t.shape,
522                        t.numeric_dtype().class_name(),
523                        t.len()
524                    );
525                    let dims = match t.integer_storage() {
526                        Some(storage) => dims_from_integer_tensor_values(storage, &t.shape)?,
527                        None => {
528                            dims_from_tensor_values(tensor_values_f64_cow(&t).as_ref(), &t.shape)?
529                        }
530                    };
531                    if dims.is_none() {
532                        tracing::debug!(
533                            gpu_shape = ?handle.shape,
534                            host_shape = ?t.shape,
535                            "dims_from_value_async: GPU tensor not interpretable as dims"
536                        );
537                    }
538                    Ok(dims)
539                }
540                Value::LogicalArray(la) => {
541                    let values: Vec<f64> = la
542                        .data
543                        .iter()
544                        .map(|&b| if b != 0 { 1.0 } else { 0.0 })
545                        .collect();
546                    let dims = dims_from_tensor_values(&values, &la.shape)?;
547                    if dims.is_none() {
548                        tracing::debug!(
549                            gpu_shape = ?handle.shape,
550                            host_shape = ?la.shape,
551                            "dims_from_value_async: GPU logical not interpretable as dims"
552                        );
553                    }
554                    Ok(dims)
555                }
556                Value::Num(n) => parse_numeric_dimension(n).map(|dim| Some(vec![dim])),
557                Value::Int(i) => parse_integer_shape_dimension(&i).map(|dim| Some(vec![dim])),
558                _ => Ok(None),
559            }
560        }
561        _ => Ok(None),
562    }
563}
564
565/// Convert an argument into a dimension index (1-based) if possible.
566pub fn parse_dimension(value: &Value, name: &str) -> Result<usize, String> {
567    match value {
568        Value::Int(i) => parse_integer_dimension(i, name, false),
569        Value::Tensor(tensor) if is_scalar_tensor(tensor) => {
570            if let Some(storage) = tensor.integer_storage() {
571                let value = storage.value_at(0).expect("one-element integer storage");
572                return parse_integer_dimension(&value, name, false);
573            }
574            parse_numeric_dimension_value(tensor_value_f64(tensor, 0), name, false)
575        }
576        Value::Num(n) => parse_numeric_dimension_value(*n, name, false),
577        other => Err(format!(
578            "{name}: dimension must be numeric, got {:?}",
579            other
580        )),
581    }
582}
583
584fn parse_numeric_dimension_value(
585    value: f64,
586    name: &str,
587    allow_zero: bool,
588) -> Result<usize, String> {
589    if !value.is_finite() {
590        return Err(format!("{name}: dimension must be finite"));
591    }
592    let rounded = value.round();
593    // Allow small floating error tolerance when users pass float-typed dims
594    if (rounded - value).abs() > 1e-6 {
595        return Err(format!("{name}: dimension must be an integer"));
596    }
597    let min = if allow_zero { 0.0 } else { 1.0 };
598    if rounded < min {
599        let bound = if allow_zero { 0 } else { 1 };
600        return Err(format!("{name}: dimension must be >= {bound}"));
601    }
602    if !fits_platform_usize(rounded) {
603        return Err(format!("{name}: dimension is outside the supported range"));
604    }
605    Ok(rounded as usize)
606}
607
608/// Attempt to extract a string from a runtime value.
609pub fn value_to_string(value: &Value) -> Option<String> {
610    String::try_from(value).ok()
611}
612
613/// Return a canonical 2-D shape for a tensor given its shape slice and element count.
614///
615/// * Empty data (`len == 0`) → `[0, 1]` (MATLAB convention for empty arrays).
616/// * No shape info (`shape.is_empty()`) → `[1, 1]` (scalar).
617/// * Otherwise → the tensor's own shape.
618pub fn default_shape_for(shape: &[usize], len: usize) -> Vec<usize> {
619    if len == 0 {
620        vec![0, 1]
621    } else if shape.is_empty() {
622        vec![1, 1]
623    } else {
624        shape.to_vec()
625    }
626}
627
628/// Clamp a scalar f64 to the uint8 range [0, 255], rounding to the nearest integer.
629pub fn clamp_u8(value: f64) -> f64 {
630    value.round().clamp(0.0, u8::MAX as f64)
631}
632
633/// Clamp a scalar f64 to the uint16 range [0, 65535], rounding to the nearest integer.
634pub fn clamp_u16(value: f64) -> f64 {
635    value.round().clamp(0.0, u16::MAX as f64)
636}
637
638/// Clamp a scalar f64 to the uint32 range [0, 4294967295], rounding to the nearest integer.
639pub fn clamp_u32(value: f64) -> f64 {
640    value.round().clamp(0.0, u32::MAX as f64)
641}
642
643/// Cast all elements of a tensor to the target dtype through typed constructors.
644pub fn coerce_tensor_dtype(tensor: Tensor, dtype: NumericDType) -> Tensor {
645    let shape = tensor.shape.clone();
646    let storage = tensor
647        .into_numeric_storage()
648        .expect("validated tensor storage");
649    match dtype {
650        NumericDType::F64 => {
651            Tensor::from_numeric_storage(NumericStorage::F64(storage.materialize_f64()), shape)
652                .expect("dtype coercion preserves the tensor element count")
653        }
654        NumericDType::F32 => {
655            Tensor::from_numeric_storage(NumericStorage::F32(storage.materialize_f32()), shape)
656                .expect("dtype coercion preserves the tensor element count")
657        }
658        integer_dtype => {
659            let prototype = match integer_dtype {
660                NumericDType::I8 => IntegerStorage::I8(Vec::new()),
661                NumericDType::I16 => IntegerStorage::I16(Vec::new()),
662                NumericDType::I32 => IntegerStorage::I32(Vec::new()),
663                NumericDType::I64 => IntegerStorage::I64(Vec::new()),
664                NumericDType::U8 => IntegerStorage::U8(Vec::new()),
665                NumericDType::U16 => IntegerStorage::U16(Vec::new()),
666                NumericDType::U32 => IntegerStorage::U32(Vec::new()),
667                NumericDType::U64 => IntegerStorage::U64(Vec::new()),
668                NumericDType::F32 | NumericDType::F64 => unreachable!(),
669            };
670            let floating_storage = match storage.into_integer_storage() {
671                Ok(storage) => {
672                    let values = storage
673                        .exact_values()
674                        .into_iter()
675                        .map(|value| prototype.cast_exact_assignment(&value))
676                        .collect();
677                    return Tensor::new_integer(
678                        prototype
679                            .from_same_class_values(values)
680                            .expect("integer coercion produces target-class values"),
681                        shape,
682                    )
683                    .expect("dtype coercion preserves the tensor element count");
684                }
685                Err(storage) => storage,
686            };
687            integer_tensor_from_f64_like(&prototype, floating_storage.materialize_f64(), &shape)
688                .expect("dtype coercion preserves the tensor element count")
689        }
690    }
691}
692
693#[cfg(test)]
694mod dtype_tests {
695    use super::{coerce_tensor_dtype, ones_with_dtype, zeros_with_dtype};
696    use runmat_value::{IntegerStorage, NumericDType, Tensor};
697
698    #[test]
699    fn dtype_directed_constructors_materialize_all_integer_classes() {
700        let cases = [
701            (NumericDType::I8, IntegerStorage::I8(vec![0, 0])),
702            (NumericDType::I16, IntegerStorage::I16(vec![0, 0])),
703            (NumericDType::I32, IntegerStorage::I32(vec![0, 0])),
704            (NumericDType::I64, IntegerStorage::I64(vec![0, 0])),
705            (NumericDType::U8, IntegerStorage::U8(vec![0, 0])),
706            (NumericDType::U16, IntegerStorage::U16(vec![0, 0])),
707            (NumericDType::U32, IntegerStorage::U32(vec![0, 0])),
708            (NumericDType::U64, IntegerStorage::U64(vec![0, 0])),
709        ];
710
711        for (dtype, expected_zeros) in cases {
712            let zeros = zeros_with_dtype(&[1, 2], dtype).expect("zeros");
713            assert_eq!(zeros.numeric_dtype(), dtype);
714            assert_eq!(zeros.integer_storage(), Some(&expected_zeros));
715
716            let ones = ones_with_dtype(&[1, 2], dtype).expect("ones");
717            assert_eq!(ones.numeric_dtype(), dtype);
718            assert_eq!(ones.integer_storage(), Some(&expected_zeros.ones_like(2)));
719        }
720    }
721
722    #[test]
723    fn coercion_creates_exact_storage_and_float_conversion_clears_it() {
724        let input = Tensor::new(vec![-2.4, 2.6], vec![1, 2]).expect("input");
725        let typed = coerce_tensor_dtype(input, NumericDType::I16);
726        assert_eq!(typed.numeric_dtype(), NumericDType::I16);
727        assert_eq!(
728            typed.integer_storage(),
729            Some(&IntegerStorage::I16(vec![-2, 3]))
730        );
731
732        let float = coerce_tensor_dtype(typed, NumericDType::F64);
733        assert_eq!(float.numeric_dtype(), NumericDType::F64);
734        assert!(float.integer_storage().is_none());
735    }
736
737    #[test]
738    fn integer_to_integer_coercion_reads_exact_storage_not_f64_mirror() {
739        let wide = 9_007_199_254_740_993_u64;
740        let input = Tensor::new_integer(IntegerStorage::U64(vec![wide, u64::MAX]), vec![1, 2])
741            .expect("input");
742
743        let same_class = coerce_tensor_dtype(input.clone(), NumericDType::U64);
744        assert_eq!(same_class.numeric_dtype(), NumericDType::U64);
745        assert_eq!(
746            same_class.integer_storage(),
747            Some(&IntegerStorage::U64(vec![wide, u64::MAX]))
748        );
749
750        let signed = coerce_tensor_dtype(input, NumericDType::I64);
751        assert_eq!(signed.numeric_dtype(), NumericDType::I64);
752        assert_eq!(
753            signed.integer_storage(),
754            Some(&IntegerStorage::I64(vec![
755                i64::try_from(wide).expect("wide value fits int64"),
756                i64::MAX,
757            ]))
758        );
759    }
760
761    #[test]
762    fn integer_to_integer_coercion_preserves_every_integer_class_exactly() {
763        let cases = [
764            (IntegerStorage::I8(vec![i8::MIN, i8::MAX]), NumericDType::I8),
765            (
766                IntegerStorage::I16(vec![i16::MIN, i16::MAX]),
767                NumericDType::I16,
768            ),
769            (
770                IntegerStorage::I32(vec![i32::MIN, i32::MAX]),
771                NumericDType::I32,
772            ),
773            (
774                IntegerStorage::I64(vec![i64::MIN, i64::MAX]),
775                NumericDType::I64,
776            ),
777            (IntegerStorage::U8(vec![0, u8::MAX]), NumericDType::U8),
778            (IntegerStorage::U16(vec![0, u16::MAX]), NumericDType::U16),
779            (IntegerStorage::U32(vec![0, u32::MAX]), NumericDType::U32),
780            (IntegerStorage::U64(vec![0, u64::MAX]), NumericDType::U64),
781        ];
782
783        for (storage, dtype) in cases {
784            let input = Tensor::new_integer(storage.clone(), vec![1, 2]).expect("integer input");
785            let output = coerce_tensor_dtype(input, dtype);
786            assert_eq!(output.numeric_dtype(), dtype);
787            assert_eq!(output.integer_storage(), Some(&storage));
788        }
789    }
790
791    #[test]
792    fn integer_to_integer_coercion_preserves_empty_shape_and_storage_class() {
793        let input =
794            Tensor::new_integer(IntegerStorage::I64(Vec::new()), vec![0, 3]).expect("empty input");
795        let output = coerce_tensor_dtype(input, NumericDType::U64);
796
797        assert_eq!(output.shape, vec![0, 3]);
798        assert_eq!(output.numeric_dtype(), NumericDType::U64);
799        assert_eq!(
800            output.integer_storage(),
801            Some(&IntegerStorage::U64(Vec::new()))
802        );
803    }
804}
805
806#[cfg(test)]
807mod dimension_tests {
808    use super::{
809        dimension_from_value_async, dims_from_value_async, integer_tensor_to_f64, parse_dimension,
810        scalar_f64_from_value_async, tensor_into_value, tensor_into_values_f64, tensor_values_f64,
811        tensor_values_f64_cow,
812    };
813    use futures::executor::block_on;
814    use runmat_value::{IntValue, IntegerStorage, NumericStorage, Tensor, Value};
815
816    #[test]
817    fn typed_dimension_parsers_preserve_representable_uint64_values() {
818        assert_eq!(
819            parse_dimension(&Value::Int(IntValue::U64(3)), "size"),
820            Ok(3)
821        );
822        match usize::try_from(u64::MAX) {
823            Ok(value) => assert_eq!(
824                parse_dimension(&Value::Int(IntValue::U64(u64::MAX)), "size"),
825                Ok(value)
826            ),
827            Err(_) => {
828                assert!(parse_dimension(&Value::Int(IntValue::U64(u64::MAX)), "size").is_err())
829            }
830        }
831        assert_eq!(
832            block_on(dims_from_value_async(&Value::Int(IntValue::U64(3)))),
833            Ok(Some(vec![3]))
834        );
835        assert_eq!(
836            block_on(dimension_from_value_async(
837                &Value::Int(IntValue::U64(3)),
838                "size",
839                false
840            )),
841            Ok(Some(3))
842        );
843        assert!(block_on(dims_from_value_async(&Value::Int(IntValue::I64(-1)))).is_err());
844    }
845
846    #[test]
847    fn typed_integer_tensor_dimension_parsers_use_exact_storage() {
848        let dims =
849            Tensor::new_integer(IntegerStorage::U64(vec![2, 3]), vec![1, 2]).expect("integer dims");
850        assert_eq!(
851            block_on(dims_from_value_async(&Value::Tensor(dims))),
852            Ok(Some(vec![2, 3]))
853        );
854
855        let scalar_dim = Tensor::new_integer(IntegerStorage::U64(vec![3]), vec![1, 1])
856            .expect("integer scalar dim");
857        assert_eq!(
858            block_on(dimension_from_value_async(
859                &Value::Tensor(scalar_dim),
860                "size",
861                false,
862            )),
863            Ok(Some(3))
864        );
865    }
866
867    #[test]
868    fn typed_integer_dimension_parsers_ignore_poisoned_f64_mirrors_for_all_classes() {
869        let storages = [
870            IntegerStorage::I8(vec![2]),
871            IntegerStorage::I16(vec![2]),
872            IntegerStorage::I32(vec![2]),
873            IntegerStorage::I64(vec![2]),
874            IntegerStorage::U8(vec![2]),
875            IntegerStorage::U16(vec![2]),
876            IntegerStorage::U32(vec![2]),
877            IntegerStorage::U64(vec![2]),
878        ];
879
880        for storage in storages {
881            let tensor = Tensor::new_integer(storage, vec![1, 1]).expect("integer dim");
882            assert_eq!(
883                block_on(dims_from_value_async(&Value::Tensor(tensor))),
884                Ok(Some(vec![2]))
885            );
886        }
887    }
888
889    #[test]
890    fn typed_integer_tensor_dimension_parsers_preserve_large_values_exactly() {
891        let large = 9_007_199_254_740_993_u64;
892        let scalar = Tensor::new_integer(IntegerStorage::U64(vec![large]), vec![1, 1])
893            .expect("large integer dim");
894        assert_eq!(
895            parse_dimension(&Value::Tensor(scalar.clone()), "size"),
896            Ok(large as usize)
897        );
898        assert_eq!(
899            block_on(dimension_from_value_async(
900                &Value::Tensor(scalar),
901                "size",
902                false,
903            )),
904            Ok(Some(large as usize))
905        );
906
907        let dims = Tensor::new_integer(IntegerStorage::U64(vec![large]), vec![1, 1])
908            .expect("large integer dims");
909        assert_eq!(
910            block_on(dims_from_value_async(&Value::Tensor(dims))),
911            Ok(Some(vec![large as usize]))
912        );
913    }
914
915    #[test]
916    fn typed_integer_tensor_dimension_parsers_reject_negative_values() {
917        let negative =
918            Tensor::new_integer(IntegerStorage::I64(vec![-1]), vec![1, 1]).expect("negative dim");
919        assert!(parse_dimension(&Value::Tensor(negative.clone()), "size").is_err());
920        assert!(block_on(dimension_from_value_async(
921            &Value::Tensor(negative.clone()),
922            "size",
923            false,
924        ))
925        .is_err());
926        assert!(block_on(dims_from_value_async(&Value::Tensor(negative))).is_err());
927    }
928
929    #[test]
930    fn typed_integer_tensor_f64_boundary_reads_exact_storage() {
931        let wide = 9_007_199_254_740_993_u64;
932        let scalar =
933            Tensor::new_integer(IntegerStorage::U64(vec![wide]), vec![1, 1]).expect("scalar");
934        assert_eq!(
935            block_on(scalar_f64_from_value_async(&Value::Tensor(scalar))),
936            Ok(Some(IntValue::U64(wide).to_f64()))
937        );
938
939        let tensor = Tensor::new_integer(IntegerStorage::U64(vec![wide, wide - 1]), vec![1, 2])
940            .expect("integer tensor");
941        assert_eq!(
942            tensor_values_f64(&tensor),
943            vec![
944                IntValue::U64(wide).to_f64(),
945                IntValue::U64(wide - 1).to_f64()
946            ]
947        );
948
949        let normalized = integer_tensor_to_f64(tensor).expect("normalize");
950        assert!(normalized.integer_storage().is_none());
951        assert_eq!(normalized.shape, vec![1, 2]);
952        assert_eq!(
953            normalized.materialize_f64(),
954            vec![
955                IntValue::U64(wide).to_f64(),
956                IntValue::U64(wide - 1).to_f64()
957            ]
958        );
959    }
960
961    #[test]
962    fn tensor_values_f64_cow_borrows_double_storage() {
963        let tensor = Tensor::new(vec![1.0, 2.0], vec![1, 2]).expect("tensor");
964        match tensor_values_f64_cow(&tensor) {
965            std::borrow::Cow::Borrowed(values) => assert_eq!(values, &[1.0, 2.0]),
966            std::borrow::Cow::Owned(values) => panic!("expected borrowed values, got {values:?}"),
967        }
968    }
969
970    #[test]
971    fn tensor_values_f64_cow_materializes_native_single_storage() {
972        let tensor = Tensor::from_f32(vec![0.1, f32::MAX], vec![1, 2]).expect("single tensor");
973        match tensor_values_f64_cow(&tensor) {
974            std::borrow::Cow::Owned(values) => {
975                assert_eq!(values, vec![f64::from(0.1_f32), f64::from(f32::MAX)])
976            }
977            std::borrow::Cow::Borrowed(values) => {
978                panic!("expected explicit single materialization, got {values:?}")
979            }
980        }
981    }
982
983    #[test]
984    fn tensor_into_values_f64_reads_typed_integer_storage_exactly() {
985        let wide = 9_007_199_254_740_993_u64;
986        let tensor = Tensor::new_integer(IntegerStorage::U64(vec![wide]), vec![1, 1])
987            .expect("integer tensor");
988
989        assert_eq!(
990            tensor_into_values_f64(tensor),
991            vec![IntValue::U64(wide).to_f64()]
992        );
993    }
994
995    #[test]
996    fn tensor_into_value_reads_typed_integer_scalar_storage_exactly() {
997        let wide = 9_007_199_254_740_993_u64;
998        let tensor = Tensor::new_integer(IntegerStorage::U64(vec![wide]), vec![1, 1])
999            .expect("integer tensor");
1000
1001        assert_eq!(tensor_into_value(tensor), Value::Int(IntValue::U64(wide)));
1002    }
1003
1004    #[test]
1005    fn tensor_into_value_preserves_native_single_scalar_storage() {
1006        let tensor = Tensor::from_f32(vec![0.25], vec![1, 1]).expect("single tensor");
1007        let Value::Tensor(output) = tensor_into_value(tensor) else {
1008            panic!("single scalar must retain tensor class");
1009        };
1010        assert_eq!(
1011            output.into_numeric_storage().unwrap(),
1012            NumericStorage::F32(vec![0.25])
1013        );
1014    }
1015
1016    #[test]
1017    fn binary_numeric_tensors_reads_typed_integer_scalar_storage_exactly() {
1018        let wide = 9_007_199_254_740_993_u64;
1019        let lhs = Tensor::new_integer(IntegerStorage::U64(vec![wide]), vec![1, 1])
1020            .expect("integer tensor");
1021        let rhs = Tensor::new(vec![1.0, 2.0], vec![1, 2]).expect("double tensor");
1022
1023        let (lhs_values, rhs_values, shape) =
1024            super::binary_numeric_tensors(&lhs, &rhs, "plus", "plus").expect("align");
1025        assert_eq!(lhs_values, vec![IntValue::U64(wide).to_f64(); 2]);
1026        assert_eq!(rhs_values, vec![1.0, 2.0]);
1027        assert_eq!(shape, vec![1, 2]);
1028    }
1029
1030    #[test]
1031    fn binary_numeric_tensors_reads_typed_integer_array_storage_exactly() {
1032        let lhs = Tensor::new_integer(IntegerStorage::I16(vec![-2, 3]), vec![1, 2]).expect("lhs");
1033        let rhs = Tensor::new_integer(IntegerStorage::U64(vec![4, 5]), vec![1, 2]).expect("rhs");
1034
1035        let (lhs_values, rhs_values, shape) =
1036            super::binary_numeric_tensors(&lhs, &rhs, "times", "times").expect("align");
1037        assert_eq!(lhs_values, vec![-2.0, 3.0]);
1038        assert_eq!(rhs_values, vec![4.0, 5.0]);
1039        assert_eq!(shape, vec![1, 2]);
1040    }
1041
1042    #[test]
1043    fn binary_numeric_tensors_reads_typed_integer_rhs_scalar_storage_exactly() {
1044        let lhs = Tensor::new(vec![1.0, 2.0], vec![1, 2]).expect("double tensor");
1045        let rhs =
1046            Tensor::new_integer(IntegerStorage::I64(vec![-7]), vec![1, 1]).expect("integer tensor");
1047
1048        let (lhs_values, rhs_values, shape) =
1049            super::binary_numeric_tensors(&lhs, &rhs, "minus", "minus").expect("align");
1050        assert_eq!(lhs_values, vec![1.0, 2.0]);
1051        assert_eq!(rhs_values, vec![-7.0, -7.0]);
1052        assert_eq!(shape, vec![1, 2]);
1053    }
1054
1055    #[test]
1056    fn binary_numeric_tensors_preserves_shape_mismatch_error() {
1057        let lhs = Tensor::new_integer(IntegerStorage::U8(vec![1, 2]), vec![1, 2]).expect("lhs");
1058        let rhs = Tensor::new_integer(IntegerStorage::U8(vec![1, 2]), vec![2, 1]).expect("rhs");
1059
1060        let err = super::binary_numeric_tensors(&lhs, &rhs, "plus", "plus").unwrap_err();
1061        assert!(err.message().contains("matching sizes"));
1062        assert_eq!(err.context.builtin.as_deref(), Some("plus"));
1063    }
1064
1065    #[test]
1066    fn floating_dimension_parsers_reject_values_outside_platform_range() {
1067        let out_of_range = usize::MAX as f64;
1068        assert!(parse_dimension(&Value::Num(out_of_range), "size").is_err());
1069        assert!(block_on(dimension_from_value_async(
1070            &Value::Num(out_of_range),
1071            "size",
1072            false
1073        ))
1074        .is_err());
1075        assert!(block_on(dims_from_value_async(&Value::Num(out_of_range))).is_err());
1076    }
1077}
1078
1079/// Align two numeric tensors for a binary element-wise operation with scalar broadcasting.
1080///
1081/// Returns `(lhs_data, rhs_data, output_shape)`.  If either operand is a
1082/// single element it is broadcast to the other's length.  `builtin` names the
1083/// calling builtin and is embedded in the error message when the shapes are
1084/// incompatible.
1085pub fn binary_numeric_tensors(
1086    lhs: &Tensor,
1087    rhs: &Tensor,
1088    context: &str,
1089    builtin: &str,
1090) -> crate::BuiltinResult<(Vec<f64>, Vec<f64>, Vec<usize>)> {
1091    let lhs_values = tensor_values_f64_cow(lhs);
1092    let rhs_values = tensor_values_f64_cow(rhs);
1093    let lhs_shape = default_shape_for(&lhs.shape, lhs_values.len());
1094    let rhs_shape = default_shape_for(&rhs.shape, rhs_values.len());
1095    match (lhs_values.len(), rhs_values.len()) {
1096        (1, 1) => Ok((vec![lhs_values[0]], vec![rhs_values[0]], vec![1, 1])),
1097        (1, len) => Ok((vec![lhs_values[0]; len], rhs_values.into_owned(), rhs_shape)),
1098        (len, 1) => Ok((lhs_values.into_owned(), vec![rhs_values[0]; len], lhs_shape)),
1099        (left, right) if left == right && lhs_shape == rhs_shape => {
1100            Ok((lhs_values.into_owned(), rhs_values.into_owned(), lhs_shape))
1101        }
1102        _ => Err(crate::build_runtime_error(format!(
1103            "{context}: operands must be scalar or have matching sizes"
1104        ))
1105        .with_builtin(builtin)
1106        .build()),
1107    }
1108}