Skip to main content

runmat_runtime/builtins/common/
elementwise.rs

1//! Element-wise operations for matrices and scalars
2//!
3//! This module implements language-compatible element-wise operations (.*,  ./,  .^)
4//! These operations work element-by-element on matrices and support scalar broadcasting.
5
6use crate::builtins::common::matrix::matrix_power;
7use crate::builtins::common::tensor as tensor_utils;
8use crate::builtins::math::elementwise::integer_arithmetic::{try_integer_binary, IntegerBinaryOp};
9use runmat_value::{
10    ComplexStorage, ComplexTensor, IntValue, IntegerStorage, NumericStorage, Tensor, Value,
11};
12
13fn complex_pow_scalar(base_re: f64, base_im: f64, exp_re: f64, exp_im: f64) -> (f64, f64) {
14    if base_re == 0.0 && base_im == 0.0 && exp_re == 0.0 && exp_im == 0.0 {
15        return (1.0, 0.0);
16    }
17    if base_re == 0.0 && base_im == 0.0 && exp_im == 0.0 && exp_re > 0.0 {
18        return (0.0, 0.0);
19    }
20    let r = (base_re.hypot(base_im)).max(0.0);
21    if r == 0.0 {
22        return (0.0, 0.0);
23    }
24    let theta = base_im.atan2(base_re);
25    let ln_r = r.ln();
26    let a = exp_re * ln_r - exp_im * theta;
27    let b = exp_re * theta + exp_im * ln_r;
28    let mag = a.exp();
29    (mag * b.cos(), mag * b.sin())
30}
31
32fn complex_pow_scalar_f32(base_re: f32, base_im: f32, exp_re: f32, exp_im: f32) -> (f32, f32) {
33    if base_re == 0.0 && base_im == 0.0 && exp_re == 0.0 && exp_im == 0.0 {
34        return (1.0, 0.0);
35    }
36    if base_re == 0.0 && base_im == 0.0 && exp_im == 0.0 && exp_re > 0.0 {
37        return (0.0, 0.0);
38    }
39    let radius = base_re.hypot(base_im).max(0.0);
40    if radius == 0.0 {
41        return (0.0, 0.0);
42    }
43    let theta = base_im.atan2(base_re);
44    let log_radius = radius.ln();
45    let real = exp_re * log_radius - exp_im * theta;
46    let imag = exp_re * theta + exp_im * log_radius;
47    let magnitude = real.exp();
48    (magnitude * imag.cos(), magnitude * imag.sin())
49}
50
51fn scalar_real_value(value: &Value) -> Option<f64> {
52    match value {
53        Value::Num(n) => Some(*n),
54        Value::Int(i) => Some(power_domain_scalar_from_integer(i)),
55        Value::Bool(b) => Some(if *b { 1.0 } else { 0.0 }),
56        Value::Tensor(t) if tensor_utils::is_scalar_tensor(t) => {
57            Some(tensor_utils::tensor_value_f64(t, 0))
58        }
59        _ => None,
60    }
61}
62
63fn scalar_complex_value(value: &Value) -> Option<(f64, f64)> {
64    match value {
65        Value::Complex(re, im) => Some((*re, *im)),
66        Value::ComplexTensor(t) if tensor_utils::is_scalar_complex_tensor(t) => {
67            let value = tensor_utils::complex_tensor_value_complex64(t, 0);
68            Some((value.re, value.im))
69        }
70        _ => None,
71    }
72}
73
74enum PromotedComplexTensorValues<'a> {
75    Raw(&'a [(f64, f64)]),
76    Exact(Vec<num_complex::Complex64>),
77}
78
79impl PromotedComplexTensorValues<'_> {
80    fn len(&self) -> usize {
81        match self {
82            Self::Raw(values) => values.len(),
83            Self::Exact(values) => values.len(),
84        }
85    }
86
87    fn value_at(&self, index: usize) -> (f64, f64) {
88        match self {
89            Self::Raw(values) => values[index],
90            Self::Exact(values) => {
91                let value = values[index];
92                (value.re, value.im)
93            }
94        }
95    }
96}
97
98fn promoted_complex_tensor_values(
99    tensor: &runmat_value::ComplexTensor,
100) -> PromotedComplexTensorValues<'_> {
101    if let Some(values) = tensor.as_f64_slice() {
102        PromotedComplexTensorValues::Raw(values)
103    } else {
104        PromotedComplexTensorValues::Exact(tensor_utils::complex_tensor_values_complex64(tensor))
105    }
106}
107
108fn provider_scalar_from_integer(value: &IntValue) -> f64 {
109    value.to_f64()
110}
111
112fn power_domain_scalar_from_integer(value: &IntValue) -> f64 {
113    value.to_f64()
114}
115
116fn scalar_power_value(base: &Value, exponent: &Value) -> Option<Value> {
117    let base_is_complex = matches!(base, Value::Complex(_, _) | Value::ComplexTensor(_));
118    let exp_is_complex = matches!(exponent, Value::Complex(_, _) | Value::ComplexTensor(_));
119    let base_val =
120        scalar_complex_value(base).or_else(|| scalar_real_value(base).map(|v| (v, 0.0)))?;
121    let exp_val =
122        scalar_complex_value(exponent).or_else(|| scalar_real_value(exponent).map(|v| (v, 0.0)))?;
123    let (br, bi) = base_val;
124    let (er, ei) = exp_val;
125    if base_is_complex || exp_is_complex || bi != 0.0 || ei != 0.0 {
126        let (re, im) = complex_pow_scalar(br, bi, er, ei);
127        return Some(Value::Complex(re, im));
128    }
129    let pow = br.powf(er);
130    if pow.is_nan() {
131        let (re, im) = complex_pow_scalar(br, 0.0, er, 0.0);
132        Some(Value::Complex(re, im))
133    } else {
134        Some(Value::Num(pow))
135    }
136}
137
138async fn to_host_value(v: &Value) -> Result<Value, String> {
139    match v {
140        Value::GpuTensor(h) => {
141            if runmat_accelerate_api::provider_for_handle(h).is_some() {
142                let gathered = crate::dispatcher::gather_if_needed_async(v)
143                    .await
144                    .map_err(|e| e.to_string())?;
145                Ok(gathered)
146            } else {
147                // Fallback: zeros tensor with same shape
148                let total: usize = h.shape.iter().product();
149                Ok(Value::Tensor(
150                    Tensor::new(vec![0.0; total], h.shape.clone()).map_err(|e| e.to_string())?,
151                ))
152            }
153        }
154        other => Ok(other.clone()),
155    }
156}
157
158/// Element-wise negation: -A
159/// Supports scalars and matrices
160pub fn elementwise_neg(a: &Value) -> Result<Value, String> {
161    match a {
162        Value::Num(x) => Ok(Value::Num(-x)),
163        Value::Complex(re, im) => Ok(Value::Complex(-*re, -*im)),
164        Value::Int(value) => Ok(Value::Int(negate_integer_scalar(value.clone()))),
165        Value::Bool(b) => Ok(Value::Bool(!b)), // Boolean negation
166        Value::Tensor(m) => {
167            let shape = m.shape.clone();
168            let storage = m.clone().into_numeric_storage()?;
169            let negated = match storage {
170                NumericStorage::F64(values) => {
171                    NumericStorage::F64(values.into_iter().map(|value| -value).collect())
172                }
173                NumericStorage::F32(values) => {
174                    NumericStorage::F32(values.into_iter().map(|value| -value).collect())
175                }
176                storage => NumericStorage::from_integer_storage(negate_integer_storage(
177                    &storage
178                        .into_integer_storage()
179                        .expect("non-floating numeric storage is integer"),
180                )),
181            };
182            Tensor::from_numeric_storage(negated, shape).map(Value::Tensor)
183        }
184        _ => Err(format!("Negation not supported for type: -{a:?}")),
185    }
186}
187
188fn negate_integer_scalar(value: IntValue) -> IntValue {
189    match value {
190        IntValue::I8(value) => IntValue::I8(value.saturating_neg()),
191        IntValue::I16(value) => IntValue::I16(value.saturating_neg()),
192        IntValue::I32(value) => IntValue::I32(value.saturating_neg()),
193        IntValue::I64(value) => IntValue::I64(value.saturating_neg()),
194        IntValue::U8(_) => IntValue::U8(0),
195        IntValue::U16(_) => IntValue::U16(0),
196        IntValue::U32(_) => IntValue::U32(0),
197        IntValue::U64(_) => IntValue::U64(0),
198    }
199}
200
201fn negate_integer_storage(storage: &IntegerStorage) -> IntegerStorage {
202    match storage {
203        IntegerStorage::I8(values) => {
204            IntegerStorage::I8(values.iter().map(|value| value.saturating_neg()).collect())
205        }
206        IntegerStorage::I16(values) => {
207            IntegerStorage::I16(values.iter().map(|value| value.saturating_neg()).collect())
208        }
209        IntegerStorage::I32(values) => {
210            IntegerStorage::I32(values.iter().map(|value| value.saturating_neg()).collect())
211        }
212        IntegerStorage::I64(values) => {
213            IntegerStorage::I64(values.iter().map(|value| value.saturating_neg()).collect())
214        }
215        IntegerStorage::U8(values) => IntegerStorage::U8(vec![0; values.len()]),
216        IntegerStorage::U16(values) => IntegerStorage::U16(vec![0; values.len()]),
217        IntegerStorage::U32(values) => IntegerStorage::U32(vec![0; values.len()]),
218        IntegerStorage::U64(values) => IntegerStorage::U64(vec![0; values.len()]),
219    }
220}
221
222/// Element-wise multiplication: A .* B
223/// Supports matrix-matrix, matrix-scalar, and scalar-matrix operations
224#[async_recursion::async_recursion(?Send)]
225pub async fn elementwise_mul(a: &Value, b: &Value) -> Result<Value, String> {
226    // GPU+scalar: keep on device if provider supports scalar mul
227    if let Some(p) = runmat_accelerate_api::provider() {
228        match (a, b) {
229            (Value::GpuTensor(ga), Value::Num(s)) => {
230                if let Ok(hc) = p.scalar_mul(ga, *s) {
231                    return Ok(Value::GpuTensor(hc));
232                }
233            }
234            (Value::Num(s), Value::GpuTensor(gb)) => {
235                if let Ok(hc) = p.scalar_mul(gb, *s) {
236                    return Ok(Value::GpuTensor(hc));
237                }
238            }
239            (Value::GpuTensor(ga), Value::Int(i)) => {
240                if let Ok(hc) = p.scalar_mul(ga, provider_scalar_from_integer(i)) {
241                    return Ok(Value::GpuTensor(hc));
242                }
243            }
244            (Value::Int(i), Value::GpuTensor(gb)) => {
245                if let Ok(hc) = p.scalar_mul(gb, provider_scalar_from_integer(i)) {
246                    return Ok(Value::GpuTensor(hc));
247                }
248            }
249            _ => {}
250        }
251    }
252    // If exactly one is GPU and no scalar fast-path, gather to host and recurse
253    if matches!(a, Value::GpuTensor(_)) ^ matches!(b, Value::GpuTensor(_)) {
254        let ah = to_host_value(a).await?;
255        let bh = to_host_value(b).await?;
256        return elementwise_mul(&ah, &bh).await;
257    }
258    if let Some(p) = runmat_accelerate_api::provider() {
259        if let (Value::GpuTensor(ha), Value::GpuTensor(hb)) = (a, b) {
260            if let Ok(hc) = p.elem_mul(ha, hb).await {
261                return Ok(Value::GpuTensor(hc));
262            }
263        }
264    }
265    if let Some(result) = try_integer_binary(a, b, IntegerBinaryOp::Multiply, "times")? {
266        return Ok(result);
267    }
268    match (a, b) {
269        // Complex scalars
270        (Value::Complex(ar, ai), Value::Complex(br, bi)) => {
271            Ok(Value::Complex(ar * br - ai * bi, ar * bi + ai * br))
272        }
273        (Value::Complex(ar, ai), Value::Num(s)) => Ok(Value::Complex(ar * s, ai * s)),
274        (Value::Num(s), Value::Complex(br, bi)) => Ok(Value::Complex(s * br, s * bi)),
275        // Scalar-scalar case
276        (Value::Num(x), Value::Num(y)) => Ok(Value::Num(x * y)),
277
278        // Matrix-scalar cases (broadcasting)
279        (Value::Tensor(m), Value::Num(s)) => multiply_real_tensor_scalar(m, *s),
280        (Value::Num(s), Value::Tensor(m)) => multiply_real_tensor_scalar(m, *s),
281
282        // Matrix-matrix case
283        (Value::Tensor(m1), Value::Tensor(m2)) => {
284            if m1.rows() != m2.rows() || m1.cols() != m2.cols() {
285                return Err(format!(
286                    "Matrix dimensions must agree for element-wise multiplication: {}x{} .* {}x{}",
287                    m1.rows(),
288                    m1.cols(),
289                    m2.rows(),
290                    m2.cols()
291                ));
292            }
293            multiply_real_tensors(m1, m2)
294        }
295
296        // Complex tensors
297        (Value::ComplexTensor(m1), Value::ComplexTensor(m2)) => {
298            if m1.rows != m2.rows || m1.cols != m2.cols {
299                return Err(format!(
300                    "Matrix dimensions must agree for element-wise multiplication: {}x{} .* {}x{}",
301                    m1.rows, m1.cols, m2.rows, m2.cols
302                ));
303            }
304            multiply_complex_tensors(m1, m2)
305        }
306        (Value::ComplexTensor(m), Value::Num(s)) => multiply_complex_tensor_scalar(m, *s),
307        (Value::Num(s), Value::ComplexTensor(m)) => multiply_complex_tensor_scalar(m, *s),
308
309        _ => Err(format!(
310            "Element-wise multiplication not supported for types: {a:?} .* {b:?}"
311        )),
312    }
313}
314
315fn multiply_real_tensor_scalar(tensor: &Tensor, scalar: f64) -> Result<Value, String> {
316    let shape = tensor.shape.clone();
317    let storage = tensor.clone().into_numeric_storage()?;
318    let output = match storage {
319        NumericStorage::F64(values) => {
320            NumericStorage::F64(values.into_iter().map(|value| value * scalar).collect())
321        }
322        NumericStorage::F32(values) => {
323            let scalar = scalar as f32;
324            NumericStorage::F32(values.into_iter().map(|value| value * scalar).collect())
325        }
326        _ => {
327            return Err(
328                "element-wise integer multiplication did not use the exact integer path".into(),
329            )
330        }
331    };
332    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
333}
334
335fn multiply_real_tensors(lhs: &Tensor, rhs: &Tensor) -> Result<Value, String> {
336    let shape = lhs.shape.clone();
337    let lhs = lhs.clone().into_numeric_storage()?;
338    let rhs = rhs.clone().into_numeric_storage()?;
339    let output = match (lhs, rhs) {
340        (NumericStorage::F64(lhs), NumericStorage::F64(rhs)) => NumericStorage::F64(
341            lhs.into_iter()
342                .zip(rhs)
343                .map(|(left, right)| left * right)
344                .collect(),
345        ),
346        (NumericStorage::F32(lhs), NumericStorage::F32(rhs)) => NumericStorage::F32(
347            lhs.into_iter()
348                .zip(rhs)
349                .map(|(left, right)| left * right)
350                .collect(),
351        ),
352        (NumericStorage::F32(lhs), NumericStorage::F64(rhs)) => NumericStorage::F32(
353            lhs.into_iter()
354                .zip(rhs)
355                .map(|(left, right)| (f64::from(left) * right) as f32)
356                .collect(),
357        ),
358        (NumericStorage::F64(lhs), NumericStorage::F32(rhs)) => NumericStorage::F32(
359            lhs.into_iter()
360                .zip(rhs)
361                .map(|(left, right)| (left * f64::from(right)) as f32)
362                .collect(),
363        ),
364        _ => {
365            return Err(
366                "element-wise integer multiplication did not use the exact integer path".into(),
367            )
368        }
369    };
370    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
371}
372
373fn multiply_complex_tensors(lhs: &ComplexTensor, rhs: &ComplexTensor) -> Result<Value, String> {
374    let shape = lhs.shape.clone();
375    let output = match (lhs.complex_storage(), rhs.complex_storage()) {
376        (ComplexStorage::F64(lhs), ComplexStorage::F64(rhs)) => ComplexStorage::F64(
377            lhs.iter()
378                .zip(rhs)
379                .map(|(&(ar, ai), &(br, bi))| (ar * br - ai * bi, ar * bi + ai * br))
380                .collect(),
381        ),
382        (ComplexStorage::F32(lhs), ComplexStorage::F32(rhs)) => ComplexStorage::F32(
383            lhs.iter()
384                .zip(rhs)
385                .map(|(&(ar, ai), &(br, bi))| (ar * br - ai * bi, ar * bi + ai * br))
386                .collect(),
387        ),
388        (ComplexStorage::F32(lhs), ComplexStorage::F64(rhs)) => ComplexStorage::F32(
389            lhs.iter()
390                .zip(rhs)
391                .map(|(&(ar, ai), &(br, bi))| {
392                    let ar = f64::from(ar);
393                    let ai = f64::from(ai);
394                    ((ar * br - ai * bi) as f32, (ar * bi + ai * br) as f32)
395                })
396                .collect(),
397        ),
398        (ComplexStorage::F64(lhs), ComplexStorage::F32(rhs)) => ComplexStorage::F32(
399            lhs.iter()
400                .zip(rhs)
401                .map(|(&(ar, ai), &(br, bi))| {
402                    let br = f64::from(br);
403                    let bi = f64::from(bi);
404                    ((ar * br - ai * bi) as f32, (ar * bi + ai * br) as f32)
405                })
406                .collect(),
407        ),
408        _ => return multiply_promoted_complex_tensors(lhs, rhs),
409    };
410    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
411}
412
413fn multiply_promoted_complex_tensors(
414    lhs: &ComplexTensor,
415    rhs: &ComplexTensor,
416) -> Result<Value, String> {
417    let lhs_values = promoted_complex_tensor_values(lhs);
418    let rhs_values = promoted_complex_tensor_values(rhs);
419    let mut output = Vec::with_capacity(lhs_values.len());
420    for index in 0..lhs_values.len() {
421        let (ar, ai) = lhs_values.value_at(index);
422        let (br, bi) = rhs_values.value_at(index);
423        output.push((ar * br - ai * bi, ar * bi + ai * br));
424    }
425    ComplexTensor::new(output, lhs.shape.clone()).map(Value::ComplexTensor)
426}
427
428fn multiply_complex_tensor_scalar(tensor: &ComplexTensor, scalar: f64) -> Result<Value, String> {
429    let shape = tensor.shape.clone();
430    let output = match tensor.complex_storage() {
431        ComplexStorage::F64(values) => ComplexStorage::F64(
432            values
433                .iter()
434                .map(|&(real, imag)| (real * scalar, imag * scalar))
435                .collect(),
436        ),
437        ComplexStorage::F32(values) => {
438            let scalar = scalar as f32;
439            ComplexStorage::F32(
440                values
441                    .iter()
442                    .map(|&(real, imag)| (real * scalar, imag * scalar))
443                    .collect(),
444            )
445        }
446        ComplexStorage::Integer(_) => {
447            return multiply_promoted_complex_tensor_scalar(tensor, scalar)
448        }
449    };
450    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
451}
452
453fn multiply_promoted_complex_tensor_scalar(
454    tensor: &ComplexTensor,
455    scalar: f64,
456) -> Result<Value, String> {
457    let values = promoted_complex_tensor_values(tensor);
458    let output = (0..values.len())
459        .map(|index| {
460            let (real, imag) = values.value_at(index);
461            (real * scalar, imag * scalar)
462        })
463        .collect();
464    ComplexTensor::new(output, tensor.shape.clone()).map(Value::ComplexTensor)
465}
466
467// elementwise_add has been retired in favor of the `plus` builtin
468
469// elementwise_sub has been retired in favor of the `minus` builtin
470
471/// Element-wise division: A ./ B
472/// Supports matrix-matrix, matrix-scalar, and scalar-matrix operations
473#[async_recursion::async_recursion(?Send)]
474pub async fn elementwise_div(a: &Value, b: &Value) -> Result<Value, String> {
475    // GPU+scalar: use scalar div when form is G ./ s or left-scalar s ./ G
476    if let Some(p) = runmat_accelerate_api::provider() {
477        match (a, b) {
478            (Value::GpuTensor(ga), Value::Num(s)) => {
479                if let Ok(hc) = p.scalar_div(ga, *s) {
480                    return Ok(Value::GpuTensor(hc));
481                }
482            }
483            (Value::GpuTensor(ga), Value::Int(i)) => {
484                if let Ok(hc) = p.scalar_div(ga, provider_scalar_from_integer(i)) {
485                    return Ok(Value::GpuTensor(hc));
486                }
487            }
488            (Value::Num(s), Value::GpuTensor(gb)) => {
489                if let Ok(hc) = p.scalar_rdiv(gb, *s) {
490                    return Ok(Value::GpuTensor(hc));
491                }
492            }
493            (Value::Int(i), Value::GpuTensor(gb)) => {
494                if let Ok(hc) = p.scalar_rdiv(gb, provider_scalar_from_integer(i)) {
495                    return Ok(Value::GpuTensor(hc));
496                }
497            }
498            _ => {}
499        }
500    }
501    if matches!(a, Value::GpuTensor(_)) ^ matches!(b, Value::GpuTensor(_)) {
502        let ah = to_host_value(a).await?;
503        let bh = to_host_value(b).await?;
504        return elementwise_div(&ah, &bh).await;
505    }
506    if let Some(p) = runmat_accelerate_api::provider() {
507        if let (Value::GpuTensor(ha), Value::GpuTensor(hb)) = (a, b) {
508            if let Ok(hc) = p.elem_div(ha, hb).await {
509                return Ok(Value::GpuTensor(hc));
510            }
511        }
512    }
513    if let Some(result) = try_integer_binary(a, b, IntegerBinaryOp::Divide, "rdivide")? {
514        return Ok(result);
515    }
516    match (a, b) {
517        // Complex scalars
518        (Value::Complex(ar, ai), Value::Complex(br, bi)) => {
519            let denom = br * br + bi * bi;
520            if denom == 0.0 {
521                return Ok(Value::Num(f64::NAN));
522            }
523            Ok(Value::Complex(
524                (ar * br + ai * bi) / denom,
525                (ai * br - ar * bi) / denom,
526            ))
527        }
528        (Value::Complex(ar, ai), Value::Num(s)) => Ok(Value::Complex(ar / s, ai / s)),
529        (Value::Num(s), Value::Complex(br, bi)) => {
530            let denom = br * br + bi * bi;
531            if denom == 0.0 {
532                return Ok(Value::Num(f64::NAN));
533            }
534            Ok(Value::Complex((s * br) / denom, (-s * bi) / denom))
535        }
536        // Scalar-scalar case
537        (Value::Num(x), Value::Num(y)) => {
538            if *y == 0.0 {
539                Ok(Value::Num(f64::INFINITY * x.signum()))
540            } else {
541                Ok(Value::Num(x / y))
542            }
543        }
544        // Matrix-scalar cases (broadcasting)
545        (Value::Tensor(m), Value::Num(s)) => divide_real_tensor_scalar(m, *s),
546        (Value::Num(s), Value::Tensor(m)) => divide_scalar_real_tensor(*s, m),
547
548        // Matrix-matrix case
549        (Value::Tensor(m1), Value::Tensor(m2)) => {
550            if m1.rows() != m2.rows() || m1.cols() != m2.cols() {
551                return Err(format!(
552                    "Matrix dimensions must agree for element-wise division: {}x{} ./ {}x{}",
553                    m1.rows(),
554                    m1.cols(),
555                    m2.rows(),
556                    m2.cols()
557                ));
558            }
559            divide_real_tensors(m1, m2)
560        }
561
562        // Complex tensors
563        (Value::ComplexTensor(m1), Value::ComplexTensor(m2)) => {
564            if m1.rows != m2.rows || m1.cols != m2.cols {
565                return Err(format!(
566                    "Matrix dimensions must agree for element-wise division: {}x{} ./ {}x{}",
567                    m1.rows, m1.cols, m2.rows, m2.cols
568                ));
569            }
570            divide_complex_tensors(m1, m2)
571        }
572        (Value::ComplexTensor(m), Value::Num(s)) => divide_complex_tensor_scalar(m, *s),
573        (Value::Num(s), Value::ComplexTensor(m)) => divide_scalar_complex_tensor(*s, m),
574
575        _ => Err(format!(
576            "Element-wise division not supported for types: {a:?} ./ {b:?}"
577        )),
578    }
579}
580
581fn divide_real_tensor_scalar(tensor: &Tensor, scalar: f64) -> Result<Value, String> {
582    let shape = tensor.shape.clone();
583    let storage = tensor.clone().into_numeric_storage()?;
584    let output = match storage {
585        NumericStorage::F64(values) => NumericStorage::F64(
586            values
587                .into_iter()
588                .map(|value| divide_real_value_f64(value, scalar))
589                .collect(),
590        ),
591        NumericStorage::F32(values) => {
592            let scalar = scalar as f32;
593            NumericStorage::F32(
594                values
595                    .into_iter()
596                    .map(|value| divide_real_value_f32(value, scalar))
597                    .collect(),
598            )
599        }
600        _ => return Err("element-wise integer division did not use the exact integer path".into()),
601    };
602    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
603}
604
605fn divide_scalar_real_tensor(scalar: f64, tensor: &Tensor) -> Result<Value, String> {
606    let shape = tensor.shape.clone();
607    let storage = tensor.clone().into_numeric_storage()?;
608    let output = match storage {
609        NumericStorage::F64(values) => NumericStorage::F64(
610            values
611                .into_iter()
612                .map(|value| divide_real_value_f64(scalar, value))
613                .collect(),
614        ),
615        NumericStorage::F32(values) => {
616            let scalar = scalar as f32;
617            NumericStorage::F32(
618                values
619                    .into_iter()
620                    .map(|value| divide_real_value_f32(scalar, value))
621                    .collect(),
622            )
623        }
624        _ => return Err("element-wise integer division did not use the exact integer path".into()),
625    };
626    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
627}
628
629fn divide_real_tensors(lhs: &Tensor, rhs: &Tensor) -> Result<Value, String> {
630    let shape = lhs.shape.clone();
631    let lhs = lhs.clone().into_numeric_storage()?;
632    let rhs = rhs.clone().into_numeric_storage()?;
633    let output = match (lhs, rhs) {
634        (NumericStorage::F64(lhs), NumericStorage::F64(rhs)) => NumericStorage::F64(
635            lhs.into_iter()
636                .zip(rhs)
637                .map(|(left, right)| divide_real_value_f64(left, right))
638                .collect(),
639        ),
640        (NumericStorage::F32(lhs), NumericStorage::F32(rhs)) => NumericStorage::F32(
641            lhs.into_iter()
642                .zip(rhs)
643                .map(|(left, right)| divide_real_value_f32(left, right))
644                .collect(),
645        ),
646        (NumericStorage::F32(lhs), NumericStorage::F64(rhs)) => NumericStorage::F32(
647            lhs.into_iter()
648                .zip(rhs)
649                .map(|(left, right)| divide_real_value_f64(f64::from(left), right) as f32)
650                .collect(),
651        ),
652        (NumericStorage::F64(lhs), NumericStorage::F32(rhs)) => NumericStorage::F32(
653            lhs.into_iter()
654                .zip(rhs)
655                .map(|(left, right)| divide_real_value_f64(left, f64::from(right)) as f32)
656                .collect(),
657        ),
658        _ => return Err("element-wise integer division did not use the exact integer path".into()),
659    };
660    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
661}
662
663fn divide_real_value_f64(numerator: f64, denominator: f64) -> f64 {
664    if denominator == 0.0 {
665        f64::INFINITY * numerator.signum()
666    } else {
667        numerator / denominator
668    }
669}
670
671fn divide_real_value_f32(numerator: f32, denominator: f32) -> f32 {
672    if denominator == 0.0 {
673        f32::INFINITY * numerator.signum()
674    } else {
675        numerator / denominator
676    }
677}
678
679fn divide_complex_tensors(lhs: &ComplexTensor, rhs: &ComplexTensor) -> Result<Value, String> {
680    let shape = lhs.shape.clone();
681    let output = match (lhs.complex_storage(), rhs.complex_storage()) {
682        (ComplexStorage::F64(lhs), ComplexStorage::F64(rhs)) => ComplexStorage::F64(
683            lhs.iter()
684                .zip(rhs)
685                .map(|(&left, &right)| divide_complex_value_f64(left, right))
686                .collect(),
687        ),
688        (ComplexStorage::F32(lhs), ComplexStorage::F32(rhs)) => ComplexStorage::F32(
689            lhs.iter()
690                .zip(rhs)
691                .map(|(&left, &right)| divide_complex_value_f32(left, right))
692                .collect(),
693        ),
694        (ComplexStorage::F32(lhs), ComplexStorage::F64(rhs)) => ComplexStorage::F32(
695            lhs.iter()
696                .zip(rhs)
697                .map(|(&(ar, ai), &right)| {
698                    let (real, imag) =
699                        divide_complex_value_f64((f64::from(ar), f64::from(ai)), right);
700                    (real as f32, imag as f32)
701                })
702                .collect(),
703        ),
704        (ComplexStorage::F64(lhs), ComplexStorage::F32(rhs)) => ComplexStorage::F32(
705            lhs.iter()
706                .zip(rhs)
707                .map(|(&left, &(br, bi))| {
708                    let (real, imag) =
709                        divide_complex_value_f64(left, (f64::from(br), f64::from(bi)));
710                    (real as f32, imag as f32)
711                })
712                .collect(),
713        ),
714        _ => return divide_promoted_complex_tensors(lhs, rhs),
715    };
716    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
717}
718
719fn divide_promoted_complex_tensors(
720    lhs: &ComplexTensor,
721    rhs: &ComplexTensor,
722) -> Result<Value, String> {
723    let lhs_values = promoted_complex_tensor_values(lhs);
724    let rhs_values = promoted_complex_tensor_values(rhs);
725    let output = (0..lhs_values.len())
726        .map(|index| {
727            divide_complex_value_f64(lhs_values.value_at(index), rhs_values.value_at(index))
728        })
729        .collect();
730    ComplexTensor::new(output, lhs.shape.clone()).map(Value::ComplexTensor)
731}
732
733fn divide_complex_tensor_scalar(tensor: &ComplexTensor, scalar: f64) -> Result<Value, String> {
734    let shape = tensor.shape.clone();
735    let output = match tensor.complex_storage() {
736        ComplexStorage::F64(values) => ComplexStorage::F64(
737            values
738                .iter()
739                .map(|&(real, imag)| (real / scalar, imag / scalar))
740                .collect(),
741        ),
742        ComplexStorage::F32(values) => {
743            let scalar = scalar as f32;
744            ComplexStorage::F32(
745                values
746                    .iter()
747                    .map(|&(real, imag)| (real / scalar, imag / scalar))
748                    .collect(),
749            )
750        }
751        ComplexStorage::Integer(_) => return divide_promoted_complex_tensor_scalar(tensor, scalar),
752    };
753    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
754}
755
756fn divide_promoted_complex_tensor_scalar(
757    tensor: &ComplexTensor,
758    scalar: f64,
759) -> Result<Value, String> {
760    let values = promoted_complex_tensor_values(tensor);
761    let output = (0..values.len())
762        .map(|index| {
763            let (real, imag) = values.value_at(index);
764            (real / scalar, imag / scalar)
765        })
766        .collect();
767    ComplexTensor::new(output, tensor.shape.clone()).map(Value::ComplexTensor)
768}
769
770fn divide_scalar_complex_tensor(scalar: f64, tensor: &ComplexTensor) -> Result<Value, String> {
771    let shape = tensor.shape.clone();
772    let output = match tensor.complex_storage() {
773        ComplexStorage::F64(values) => ComplexStorage::F64(
774            values
775                .iter()
776                .map(|&denominator| divide_complex_value_f64((scalar, 0.0), denominator))
777                .collect(),
778        ),
779        ComplexStorage::F32(values) => {
780            let scalar = scalar as f32;
781            ComplexStorage::F32(
782                values
783                    .iter()
784                    .map(|&denominator| divide_complex_value_f32((scalar, 0.0), denominator))
785                    .collect(),
786            )
787        }
788        ComplexStorage::Integer(_) => return divide_scalar_promoted_complex_tensor(scalar, tensor),
789    };
790    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
791}
792
793fn divide_scalar_promoted_complex_tensor(
794    scalar: f64,
795    tensor: &ComplexTensor,
796) -> Result<Value, String> {
797    let values = promoted_complex_tensor_values(tensor);
798    let output = (0..values.len())
799        .map(|index| divide_complex_value_f64((scalar, 0.0), values.value_at(index)))
800        .collect();
801    ComplexTensor::new(output, tensor.shape.clone()).map(Value::ComplexTensor)
802}
803
804fn divide_complex_value_f64(numerator: (f64, f64), denominator: (f64, f64)) -> (f64, f64) {
805    let divisor = denominator.0 * denominator.0 + denominator.1 * denominator.1;
806    if divisor == 0.0 {
807        (f64::NAN, f64::NAN)
808    } else {
809        (
810            (numerator.0 * denominator.0 + numerator.1 * denominator.1) / divisor,
811            (numerator.1 * denominator.0 - numerator.0 * denominator.1) / divisor,
812        )
813    }
814}
815
816fn divide_complex_value_f32(numerator: (f32, f32), denominator: (f32, f32)) -> (f32, f32) {
817    let divisor = denominator.0 * denominator.0 + denominator.1 * denominator.1;
818    if divisor == 0.0 {
819        (f32::NAN, f32::NAN)
820    } else {
821        (
822            (numerator.0 * denominator.0 + numerator.1 * denominator.1) / divisor,
823            (numerator.1 * denominator.0 - numerator.0 * denominator.1) / divisor,
824        )
825    }
826}
827
828/// Regular power operation: A ^ B  
829/// For matrices, this is matrix exponentiation (A^n where n is integer)
830/// For scalars, this is regular exponentiation
831pub fn power(a: &Value, b: &Value) -> Result<Value, String> {
832    if scalar_power_integer_candidate(a) && scalar_power_integer_candidate(b) {
833        if let Some(result) = try_integer_binary(a, b, IntegerBinaryOp::Power, "power")? {
834            return Ok(result);
835        }
836    }
837    if let Some(result) = scalar_power_value(a, b) {
838        return Ok(result);
839    }
840    match (a, b) {
841        // Scalar cases - include complex
842        (Value::Complex(ar, ai), Value::Complex(br, bi)) => {
843            let (r, i) = complex_pow_scalar(*ar, *ai, *br, *bi);
844            Ok(Value::Complex(r, i))
845        }
846        (Value::Complex(ar, ai), Value::Num(y)) => {
847            let (r, i) = complex_pow_scalar(*ar, *ai, *y, 0.0);
848            Ok(Value::Complex(r, i))
849        }
850        (Value::Num(x), Value::Complex(br, bi)) => {
851            let (r, i) = complex_pow_scalar(*x, 0.0, *br, *bi);
852            Ok(Value::Complex(r, i))
853        }
854        // Scalar cases - real only
855        (Value::Num(x), Value::Num(y)) => Ok(Value::Num(x.powf(*y))),
856
857        // Matrix^scalar case - matrix exponentiation
858        (Value::Tensor(m), Value::Num(s)) => {
859            let result = matrix_power(m, matrix_power_exponent_from_f64(*s)?)?;
860            Ok(Value::Tensor(result))
861        }
862        (Value::Tensor(m), Value::Int(s)) => {
863            let result = matrix_power(m, matrix_power_exponent_from_int(s)?)?;
864            Ok(Value::Tensor(result))
865        }
866
867        // Complex matrix^integer case
868        (Value::ComplexTensor(m), Value::Num(s)) => {
869            let result = crate::builtins::common::matrix::complex_matrix_power(
870                m,
871                matrix_power_exponent_from_f64(*s)?,
872            )?;
873            Ok(Value::ComplexTensor(result))
874        }
875        (Value::ComplexTensor(m), Value::Int(s)) => {
876            let result = crate::builtins::common::matrix::complex_matrix_power(
877                m,
878                matrix_power_exponent_from_int(s)?,
879            )?;
880            Ok(Value::ComplexTensor(result))
881        }
882
883        // Other cases not supported for regular matrix power
884        _ => Err(format!(
885            "Power operation not supported for types: {a:?} ^ {b:?}"
886        )),
887    }
888}
889
890fn scalar_power_integer_candidate(value: &Value) -> bool {
891    match value {
892        Value::Int(_) | Value::Num(_) | Value::Bool(_) => true,
893        Value::Tensor(tensor) => tensor_utils::is_scalar_tensor(tensor),
894        Value::LogicalArray(array) => array.data.len() == 1,
895        _ => false,
896    }
897}
898
899fn matrix_power_exponent_from_f64(value: f64) -> Result<i32, String> {
900    if !value.is_finite() || value.fract() != 0.0 {
901        return Err("Matrix power requires integer exponent".to_string());
902    }
903    if value < i32::MIN as f64 || value > i32::MAX as f64 {
904        return Err("Matrix power exponent is outside the supported int32 range".to_string());
905    }
906    Ok(value as i32)
907}
908
909fn matrix_power_exponent_from_int(value: &IntValue) -> Result<i32, String> {
910    value
911        .try_to_i32()
912        .ok_or_else(|| "Matrix power exponent is outside the supported int32 range".to_string())
913}
914
915/// Element-wise power: A .^ B
916/// Supports matrix-matrix, matrix-scalar, and scalar-matrix operations
917pub fn elementwise_pow(a: &Value, b: &Value) -> Result<Value, String> {
918    if let Some(result) = try_integer_binary(a, b, IntegerBinaryOp::Power, "power")? {
919        return Ok(result);
920    }
921    match (a, b) {
922        // Complex scalar cases
923        (Value::Complex(ar, ai), Value::Complex(br, bi)) => {
924            let (r, i) = complex_pow_scalar(*ar, *ai, *br, *bi);
925            Ok(Value::Complex(r, i))
926        }
927        (Value::Complex(ar, ai), Value::Num(y)) => {
928            let (r, i) = complex_pow_scalar(*ar, *ai, *y, 0.0);
929            Ok(Value::Complex(r, i))
930        }
931        (Value::Num(x), Value::Complex(br, bi)) => {
932            let (r, i) = complex_pow_scalar(*x, 0.0, *br, *bi);
933            Ok(Value::Complex(r, i))
934        }
935        // Scalar-scalar case
936        (Value::Num(x), Value::Num(y)) => Ok(Value::Num(x.powf(*y))),
937
938        // Matrix-scalar cases (broadcasting)
939        (Value::Tensor(m), Value::Num(s)) => power_real_tensor_scalar(m, *s),
940        (Value::Num(s), Value::Tensor(m)) => power_scalar_real_tensor(*s, m),
941
942        // Matrix-matrix case
943        (Value::Tensor(m1), Value::Tensor(m2)) => {
944            if m1.rows() != m2.rows() || m1.cols() != m2.cols() {
945                return Err(format!(
946                    "Matrix dimensions must agree for element-wise power: {}x{} .^ {}x{}",
947                    m1.rows(),
948                    m1.cols(),
949                    m2.rows(),
950                    m2.cols()
951                ));
952            }
953            power_real_tensors(m1, m2)
954        }
955
956        // Complex tensor element-wise power
957        (Value::ComplexTensor(m1), Value::ComplexTensor(m2)) => {
958            if m1.rows != m2.rows || m1.cols != m2.cols {
959                return Err(format!(
960                    "Matrix dimensions must agree for element-wise power: {}x{} .^ {}x{}",
961                    m1.rows, m1.cols, m2.rows, m2.cols
962                ));
963            }
964            power_complex_tensors(m1, m2)
965        }
966        (Value::ComplexTensor(m), Value::Num(s)) => power_complex_tensor_scalar(m, (*s, 0.0)),
967        (Value::ComplexTensor(m), Value::Complex(br, bi)) => {
968            power_complex_tensor_scalar(m, (*br, *bi))
969        }
970        (Value::Num(s), Value::ComplexTensor(m)) => power_scalar_complex_tensor((*s, 0.0), m),
971        (Value::Complex(br, bi), Value::ComplexTensor(m)) => {
972            power_scalar_complex_tensor((*br, *bi), m)
973        }
974
975        _ => Err(format!(
976            "Element-wise power not supported for types: {a:?} .^ {b:?}"
977        )),
978    }
979}
980
981fn power_real_tensor_scalar(tensor: &Tensor, exponent: f64) -> Result<Value, String> {
982    let shape = tensor.shape.clone();
983    let storage = tensor.clone().into_numeric_storage()?;
984    let output = match storage {
985        NumericStorage::F64(values) => NumericStorage::F64(
986            values
987                .into_iter()
988                .map(|value| value.powf(exponent))
989                .collect(),
990        ),
991        NumericStorage::F32(values) => {
992            let exponent = exponent as f32;
993            NumericStorage::F32(
994                values
995                    .into_iter()
996                    .map(|value| value.powf(exponent))
997                    .collect(),
998            )
999        }
1000        _ => return Err("element-wise integer power did not use the exact integer path".into()),
1001    };
1002    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
1003}
1004
1005fn power_scalar_real_tensor(base: f64, tensor: &Tensor) -> Result<Value, String> {
1006    let shape = tensor.shape.clone();
1007    let storage = tensor.clone().into_numeric_storage()?;
1008    let output = match storage {
1009        NumericStorage::F64(values) => {
1010            NumericStorage::F64(values.into_iter().map(|value| base.powf(value)).collect())
1011        }
1012        NumericStorage::F32(values) => {
1013            let base = base as f32;
1014            NumericStorage::F32(values.into_iter().map(|value| base.powf(value)).collect())
1015        }
1016        _ => return Err("element-wise integer power did not use the exact integer path".into()),
1017    };
1018    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
1019}
1020
1021fn power_real_tensors(base: &Tensor, exponent: &Tensor) -> Result<Value, String> {
1022    let shape = base.shape.clone();
1023    let base = base.clone().into_numeric_storage()?;
1024    let exponent = exponent.clone().into_numeric_storage()?;
1025    let output = match (base, exponent) {
1026        (NumericStorage::F64(base), NumericStorage::F64(exponent)) => NumericStorage::F64(
1027            base.into_iter()
1028                .zip(exponent)
1029                .map(|(base, exponent)| base.powf(exponent))
1030                .collect(),
1031        ),
1032        (NumericStorage::F32(base), NumericStorage::F32(exponent)) => NumericStorage::F32(
1033            base.into_iter()
1034                .zip(exponent)
1035                .map(|(base, exponent)| base.powf(exponent))
1036                .collect(),
1037        ),
1038        (NumericStorage::F32(base), NumericStorage::F64(exponent)) => NumericStorage::F32(
1039            base.into_iter()
1040                .zip(exponent)
1041                .map(|(base, exponent)| f64::from(base).powf(exponent) as f32)
1042                .collect(),
1043        ),
1044        (NumericStorage::F64(base), NumericStorage::F32(exponent)) => NumericStorage::F32(
1045            base.into_iter()
1046                .zip(exponent)
1047                .map(|(base, exponent)| base.powf(f64::from(exponent)) as f32)
1048                .collect(),
1049        ),
1050        _ => return Err("element-wise integer power did not use the exact integer path".into()),
1051    };
1052    Tensor::from_numeric_storage(output, shape).map(Value::Tensor)
1053}
1054
1055fn power_complex_tensors(base: &ComplexTensor, exponent: &ComplexTensor) -> Result<Value, String> {
1056    let shape = base.shape.clone();
1057    let output = match (base.complex_storage(), exponent.complex_storage()) {
1058        (ComplexStorage::F64(base), ComplexStorage::F64(exponent)) => ComplexStorage::F64(
1059            base.iter()
1060                .zip(exponent)
1061                .map(|(&(br, bi), &(er, ei))| complex_pow_scalar(br, bi, er, ei))
1062                .collect(),
1063        ),
1064        (ComplexStorage::F32(base), ComplexStorage::F32(exponent)) => ComplexStorage::F32(
1065            base.iter()
1066                .zip(exponent)
1067                .map(|(&(br, bi), &(er, ei))| complex_pow_scalar_f32(br, bi, er, ei))
1068                .collect(),
1069        ),
1070        (ComplexStorage::F32(base), ComplexStorage::F64(exponent)) => ComplexStorage::F32(
1071            base.iter()
1072                .zip(exponent)
1073                .map(|(&(br, bi), &(er, ei))| {
1074                    let (real, imag) = complex_pow_scalar(f64::from(br), f64::from(bi), er, ei);
1075                    (real as f32, imag as f32)
1076                })
1077                .collect(),
1078        ),
1079        (ComplexStorage::F64(base), ComplexStorage::F32(exponent)) => ComplexStorage::F32(
1080            base.iter()
1081                .zip(exponent)
1082                .map(|(&(br, bi), &(er, ei))| {
1083                    let (real, imag) = complex_pow_scalar(br, bi, f64::from(er), f64::from(ei));
1084                    (real as f32, imag as f32)
1085                })
1086                .collect(),
1087        ),
1088        _ => return power_promoted_complex_tensors(base, exponent),
1089    };
1090    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
1091}
1092
1093fn power_promoted_complex_tensors(
1094    base: &ComplexTensor,
1095    exponent: &ComplexTensor,
1096) -> Result<Value, String> {
1097    let base_values = promoted_complex_tensor_values(base);
1098    let exponent_values = promoted_complex_tensor_values(exponent);
1099    let output = (0..base_values.len())
1100        .map(|index| {
1101            let (br, bi) = base_values.value_at(index);
1102            let (er, ei) = exponent_values.value_at(index);
1103            complex_pow_scalar(br, bi, er, ei)
1104        })
1105        .collect();
1106    ComplexTensor::new(output, base.shape.clone()).map(Value::ComplexTensor)
1107}
1108
1109fn power_complex_tensor_scalar(
1110    base: &ComplexTensor,
1111    exponent: (f64, f64),
1112) -> Result<Value, String> {
1113    let shape = base.shape.clone();
1114    let output = match base.complex_storage() {
1115        ComplexStorage::F64(values) => ComplexStorage::F64(
1116            values
1117                .iter()
1118                .map(|&(br, bi)| complex_pow_scalar(br, bi, exponent.0, exponent.1))
1119                .collect(),
1120        ),
1121        ComplexStorage::F32(values) => {
1122            let exponent = (exponent.0 as f32, exponent.1 as f32);
1123            ComplexStorage::F32(
1124                values
1125                    .iter()
1126                    .map(|&(br, bi)| complex_pow_scalar_f32(br, bi, exponent.0, exponent.1))
1127                    .collect(),
1128            )
1129        }
1130        ComplexStorage::Integer(_) => return power_promoted_complex_tensor_scalar(base, exponent),
1131    };
1132    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
1133}
1134
1135fn power_promoted_complex_tensor_scalar(
1136    base: &ComplexTensor,
1137    exponent: (f64, f64),
1138) -> Result<Value, String> {
1139    let values = promoted_complex_tensor_values(base);
1140    let output = (0..values.len())
1141        .map(|index| {
1142            let (br, bi) = values.value_at(index);
1143            complex_pow_scalar(br, bi, exponent.0, exponent.1)
1144        })
1145        .collect();
1146    ComplexTensor::new(output, base.shape.clone()).map(Value::ComplexTensor)
1147}
1148
1149fn power_scalar_complex_tensor(
1150    base: (f64, f64),
1151    exponent: &ComplexTensor,
1152) -> Result<Value, String> {
1153    let shape = exponent.shape.clone();
1154    let output = match exponent.complex_storage() {
1155        ComplexStorage::F64(values) => ComplexStorage::F64(
1156            values
1157                .iter()
1158                .map(|&(er, ei)| complex_pow_scalar(base.0, base.1, er, ei))
1159                .collect(),
1160        ),
1161        ComplexStorage::F32(values) => {
1162            let base = (base.0 as f32, base.1 as f32);
1163            ComplexStorage::F32(
1164                values
1165                    .iter()
1166                    .map(|&(er, ei)| complex_pow_scalar_f32(base.0, base.1, er, ei))
1167                    .collect(),
1168            )
1169        }
1170        ComplexStorage::Integer(_) => return power_scalar_promoted_complex_tensor(base, exponent),
1171    };
1172    ComplexTensor::from_complex_storage(output, shape).map(Value::ComplexTensor)
1173}
1174
1175fn power_scalar_promoted_complex_tensor(
1176    base: (f64, f64),
1177    exponent: &ComplexTensor,
1178) -> Result<Value, String> {
1179    let values = promoted_complex_tensor_values(exponent);
1180    let output = (0..values.len())
1181        .map(|index| {
1182            let (er, ei) = values.value_at(index);
1183            complex_pow_scalar(base.0, base.1, er, ei)
1184        })
1185        .collect();
1186    ComplexTensor::new(output, exponent.shape.clone()).map(Value::ComplexTensor)
1187}
1188
1189// Element-wise operations are not directly exposed as runtime builtins because they need
1190// to handle multiple types (Value enum variants). Instead, they are called directly from
1191// the interpreter and JIT compiler using the elementwise_* functions above.
1192
1193#[cfg(test)]
1194mod tests {
1195    use super::*;
1196    use futures::executor::block_on;
1197
1198    #[test]
1199    fn matrix_power_typed_exponent_parser_is_exact() {
1200        assert_eq!(
1201            matrix_power_exponent_from_int(&IntValue::U16(7)).unwrap(),
1202            7
1203        );
1204        assert!(matrix_power_exponent_from_int(&IntValue::U64(u64::MAX)).is_err());
1205        assert!(matrix_power_exponent_from_f64(f64::INFINITY).is_err());
1206        assert!(matrix_power_exponent_from_f64(i32::MAX as f64 + 1.0).is_err());
1207    }
1208
1209    #[test]
1210    fn scalar_power_reads_typed_complex_integer_storage_exactly() {
1211        let storage = runmat_value::IntegerComplexStorage::new(
1212            IntegerStorage::I16(vec![3]),
1213            IntegerStorage::I16(vec![4]),
1214        )
1215        .expect("complex integer storage");
1216        let tensor =
1217            runmat_value::ComplexTensor::new_integer(storage, vec![1, 1]).expect("complex tensor");
1218
1219        let result = scalar_power_value(&Value::ComplexTensor(tensor), &Value::Num(1.0))
1220            .expect("scalar power");
1221        match result {
1222            Value::Complex(re, im) => {
1223                assert!((re - 3.0).abs() < 1e-12);
1224                assert!((im - 4.0).abs() < 1e-12);
1225            }
1226            other => panic!("expected complex scalar, got {other:?}"),
1227        }
1228    }
1229
1230    fn mirrorless_complex_integer_tensor(
1231        real: Vec<i16>,
1232        imag: Vec<i16>,
1233        shape: Vec<usize>,
1234    ) -> runmat_value::ComplexTensor {
1235        let storage = runmat_value::IntegerComplexStorage::new(
1236            IntegerStorage::I16(real),
1237            IntegerStorage::I16(imag),
1238        )
1239        .expect("complex integer storage");
1240
1241        runmat_value::ComplexTensor::new_integer(storage, shape).expect("complex tensor")
1242    }
1243
1244    #[test]
1245    fn elementwise_mul_reads_typed_complex_integer_storage_exactly() {
1246        let lhs = mirrorless_complex_integer_tensor(vec![3, -2], vec![4, 5], vec![1, 2]);
1247        let rhs = mirrorless_complex_integer_tensor(vec![1, 6], vec![-2, 1], vec![1, 2]);
1248
1249        let Value::ComplexTensor(result) = block_on(elementwise_mul(
1250            &Value::ComplexTensor(lhs),
1251            &Value::ComplexTensor(rhs),
1252        ))
1253        .expect("mul") else {
1254            panic!("expected complex tensor");
1255        };
1256        assert_eq!(result.shape, vec![1, 2]);
1257        assert_eq!(result.materialize_f64(), vec![(11.0, -2.0), (-17.0, 28.0)]);
1258    }
1259
1260    #[test]
1261    fn elementwise_mul_preserves_native_complex_single_and_nd_shape() {
1262        let shape = vec![1, 2, 1];
1263        let lhs = ComplexTensor::from_f32(vec![(1.0, 2.0), (-2.0, 1.0)], shape.clone()).unwrap();
1264        let rhs = ComplexTensor::from_f32(vec![(3.0, -1.0), (4.0, 2.0)], shape.clone()).unwrap();
1265        let Value::ComplexTensor(result) = block_on(elementwise_mul(
1266            &Value::ComplexTensor(lhs),
1267            &Value::ComplexTensor(rhs),
1268        ))
1269        .expect("mul") else {
1270            panic!("expected complex tensor");
1271        };
1272        assert_eq!(result.shape, shape);
1273        assert_eq!(result.as_f32_slice(), Some(&[(5.0, 5.0), (-10.0, 0.0)][..]));
1274    }
1275
1276    #[test]
1277    fn elementwise_mul_mixed_complex_floating_returns_single() {
1278        let single = ComplexTensor::from_f32(vec![(1.0, 2.0)], vec![1, 1]).unwrap();
1279        let double = ComplexTensor::new(vec![(3.0, -1.0)], vec![1, 1]).unwrap();
1280        for (lhs, rhs) in [
1281            (single.clone(), double.clone()),
1282            (double.clone(), single.clone()),
1283        ] {
1284            let Value::ComplexTensor(result) = block_on(elementwise_mul(
1285                &Value::ComplexTensor(lhs),
1286                &Value::ComplexTensor(rhs),
1287            ))
1288            .expect("mul") else {
1289                panic!("expected complex tensor");
1290            };
1291            assert_eq!(result.as_f32_slice(), Some(&[(5.0, 5.0)][..]));
1292        }
1293    }
1294
1295    #[test]
1296    fn elementwise_mul_complex_single_by_double_scalar_returns_single() {
1297        let tensor = ComplexTensor::from_f32(vec![(1.0, 2.0), (-2.0, 1.0)], vec![1, 2]).unwrap();
1298        let Value::ComplexTensor(result) = block_on(elementwise_mul(
1299            &Value::ComplexTensor(tensor),
1300            &Value::Num(0.5),
1301        ))
1302        .expect("mul") else {
1303            panic!("expected complex tensor");
1304        };
1305        assert_eq!(result.as_f32_slice(), Some(&[(0.5, 1.0), (-1.0, 0.5)][..]));
1306    }
1307
1308    #[test]
1309    fn elementwise_div_reads_typed_complex_integer_storage_exactly() {
1310        let lhs = mirrorless_complex_integer_tensor(vec![3, -2], vec![4, 5], vec![1, 2]);
1311
1312        let Value::ComplexTensor(result) = block_on(elementwise_div(
1313            &Value::ComplexTensor(lhs),
1314            &Value::Num(2.0),
1315        ))
1316        .expect("div") else {
1317            panic!("expected complex tensor");
1318        };
1319        assert_eq!(result.shape, vec![1, 2]);
1320        assert_eq!(result.materialize_f64(), vec![(1.5, 2.0), (-1.0, 2.5)]);
1321    }
1322
1323    #[test]
1324    fn elementwise_div_preserves_native_complex_single_and_nd_shape() {
1325        let shape = vec![1, 2, 1];
1326        let lhs = ComplexTensor::from_f32(vec![(5.0, 5.0), (-10.0, 0.0)], shape.clone()).unwrap();
1327        let rhs = ComplexTensor::from_f32(vec![(3.0, -1.0), (4.0, 2.0)], shape.clone()).unwrap();
1328        let Value::ComplexTensor(result) = block_on(elementwise_div(
1329            &Value::ComplexTensor(lhs),
1330            &Value::ComplexTensor(rhs),
1331        ))
1332        .expect("div") else {
1333            panic!("expected complex tensor");
1334        };
1335        assert_eq!(result.shape, shape);
1336        assert_eq!(result.as_f32_slice(), Some(&[(1.0, 2.0), (-2.0, 1.0)][..]));
1337    }
1338
1339    #[test]
1340    fn elementwise_div_mixed_complex_floating_returns_single() {
1341        let single = ComplexTensor::from_f32(vec![(5.0, 5.0)], vec![1, 1]).unwrap();
1342        let double = ComplexTensor::new(vec![(3.0, -1.0)], vec![1, 1]).unwrap();
1343        let Value::ComplexTensor(result) = block_on(elementwise_div(
1344            &Value::ComplexTensor(single),
1345            &Value::ComplexTensor(double),
1346        ))
1347        .expect("div") else {
1348            panic!("expected complex tensor");
1349        };
1350        assert_eq!(result.as_f32_slice(), Some(&[(1.0, 2.0)][..]));
1351    }
1352
1353    #[test]
1354    fn elementwise_div_complex_single_scalar_paths_preserve_single() {
1355        let tensor = ComplexTensor::from_f32(vec![(2.0, 4.0)], vec![1, 1]).unwrap();
1356        let Value::ComplexTensor(by_scalar) = block_on(elementwise_div(
1357            &Value::ComplexTensor(tensor.clone()),
1358            &Value::Num(2.0),
1359        ))
1360        .expect("div") else {
1361            panic!("expected complex tensor");
1362        };
1363        assert_eq!(by_scalar.as_f32_slice(), Some(&[(1.0, 2.0)][..]));
1364
1365        let Value::ComplexTensor(scalar_by) = block_on(elementwise_div(
1366            &Value::Num(10.0),
1367            &Value::ComplexTensor(tensor),
1368        ))
1369        .expect("div") else {
1370            panic!("expected complex tensor");
1371        };
1372        assert_eq!(scalar_by.as_f32_slice(), Some(&[(1.0, -2.0)][..]));
1373    }
1374
1375    #[test]
1376    fn elementwise_pow_reads_typed_complex_integer_storage_exactly() {
1377        let base = mirrorless_complex_integer_tensor(vec![3, 1], vec![4, -2], vec![1, 2]);
1378
1379        let Value::ComplexTensor(result) =
1380            elementwise_pow(&Value::ComplexTensor(base), &Value::Num(2.0)).expect("pow")
1381        else {
1382            panic!("expected complex tensor");
1383        };
1384        assert_eq!(result.shape, vec![1, 2]);
1385        assert!((result.materialize_f64()[0].0 + 7.0).abs() < 1e-12);
1386        assert!((result.materialize_f64()[0].1 - 24.0).abs() < 1e-12);
1387        assert!((result.materialize_f64()[1].0 + 3.0).abs() < 1e-12);
1388        assert!((result.materialize_f64()[1].1 + 4.0).abs() < 1e-12);
1389    }
1390
1391    #[test]
1392    fn elementwise_pow_preserves_native_complex_single_and_nd_shape() {
1393        let shape = vec![1, 2, 1];
1394        let base = ComplexTensor::from_f32(vec![(1.0, 2.0), (2.0, -1.0)], shape.clone()).unwrap();
1395        let exponent =
1396            ComplexTensor::from_f32(vec![(2.0, 0.0), (2.0, 0.0)], shape.clone()).unwrap();
1397        let Value::ComplexTensor(result) =
1398            elementwise_pow(&Value::ComplexTensor(base), &Value::ComplexTensor(exponent))
1399                .expect("power")
1400        else {
1401            panic!("expected complex tensor");
1402        };
1403        assert_eq!(result.shape, shape);
1404        let values = result.as_f32_slice().expect("single storage");
1405        for (actual, expected) in values.iter().zip([(-3.0, 4.0), (3.0, -4.0)]) {
1406            assert!((actual.0 - expected.0).abs() < 1e-5);
1407            assert!((actual.1 - expected.1).abs() < 1e-5);
1408        }
1409    }
1410
1411    #[test]
1412    fn elementwise_pow_mixed_complex_floating_returns_single() {
1413        let single = ComplexTensor::from_f32(vec![(1.0, 2.0)], vec![1, 1]).unwrap();
1414        let double = ComplexTensor::new(vec![(2.0, 0.0)], vec![1, 1]).unwrap();
1415        for (base, exponent) in [
1416            (single.clone(), double.clone()),
1417            (double.clone(), single.clone()),
1418        ] {
1419            let Value::ComplexTensor(result) =
1420                elementwise_pow(&Value::ComplexTensor(base), &Value::ComplexTensor(exponent))
1421                    .expect("power")
1422            else {
1423                panic!("expected complex tensor");
1424            };
1425            assert!(result.as_f32_slice().is_some());
1426        }
1427    }
1428
1429    #[test]
1430    fn elementwise_pow_complex_single_scalar_paths_preserve_single() {
1431        let tensor = ComplexTensor::from_f32(vec![(1.0, 2.0)], vec![1, 1]).unwrap();
1432        let Value::ComplexTensor(tensor_base) =
1433            elementwise_pow(&Value::ComplexTensor(tensor.clone()), &Value::Num(2.0))
1434                .expect("power")
1435        else {
1436            panic!("expected complex tensor");
1437        };
1438        let value = tensor_base.as_f32_slice().expect("single storage")[0];
1439        assert!((value.0 + 3.0).abs() < 1e-5);
1440        assert!((value.1 - 4.0).abs() < 1e-5);
1441
1442        let Value::ComplexTensor(tensor_exponent) =
1443            elementwise_pow(&Value::Num(2.0), &Value::ComplexTensor(tensor)).expect("power")
1444        else {
1445            panic!("expected complex tensor");
1446        };
1447        assert!(tensor_exponent.as_f32_slice().is_some());
1448    }
1449
1450    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1451    #[test]
1452    fn test_elementwise_mul_scalars() {
1453        assert_eq!(
1454            block_on(elementwise_mul(&Value::Num(3.0), &Value::Num(4.0))).unwrap(),
1455            Value::Num(12.0)
1456        );
1457        assert_eq!(
1458            block_on(elementwise_mul(
1459                &Value::Int(runmat_value::IntValue::I32(3)),
1460                &Value::Num(4.5)
1461            ))
1462            .unwrap(),
1463            Value::Int(runmat_value::IntValue::I32(14))
1464        );
1465    }
1466
1467    #[test]
1468    fn elementwise_mul_preserves_native_single_and_nd_shape() {
1469        let shape = vec![1, 2, 2];
1470        let tensor = Tensor::from_f32(vec![1.0, 2.0, 3.0, 4.0], shape.clone()).unwrap();
1471        let Value::Tensor(result) =
1472            block_on(elementwise_mul(&Value::Tensor(tensor), &Value::Num(0.5))).expect("mul")
1473        else {
1474            panic!("expected tensor");
1475        };
1476        assert_eq!(result.shape, shape);
1477        assert_eq!(
1478            result.into_numeric_storage().expect("storage"),
1479            NumericStorage::F32(vec![0.5, 1.0, 1.5, 2.0])
1480        );
1481    }
1482
1483    #[test]
1484    fn elementwise_mul_mixed_floating_tensors_returns_single() {
1485        let single = Tensor::from_f32(vec![1.5, -2.0], vec![1, 2]).unwrap();
1486        let double = Tensor::new(vec![2.0, 4.0], vec![1, 2]).unwrap();
1487        for (lhs, rhs) in [
1488            (single.clone(), double.clone()),
1489            (double.clone(), single.clone()),
1490        ] {
1491            let Value::Tensor(result) =
1492                block_on(elementwise_mul(&Value::Tensor(lhs), &Value::Tensor(rhs))).expect("mul")
1493            else {
1494                panic!("expected tensor");
1495            };
1496            assert_eq!(
1497                result.into_numeric_storage().expect("storage"),
1498                NumericStorage::F32(vec![3.0, -8.0])
1499            );
1500        }
1501    }
1502
1503    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1504    #[test]
1505    fn test_elementwise_mul_matrix_scalar() {
1506        let matrix = Tensor::new_2d(vec![1.0, 2.0, 3.0, 4.0], 2, 2).unwrap();
1507        let result = block_on(elementwise_mul(&Value::Tensor(matrix), &Value::Num(2.0))).unwrap();
1508
1509        if let Value::Tensor(m) = result {
1510            assert_eq!(m.materialize_f64(), vec![2.0, 4.0, 6.0, 8.0]);
1511            assert_eq!(m.rows(), 2);
1512            assert_eq!(m.cols(), 2);
1513        } else {
1514            panic!("Expected matrix result");
1515        }
1516    }
1517
1518    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1519    #[test]
1520    fn test_elementwise_mul_matrices() {
1521        let m1 = Tensor::new_2d(vec![1.0, 2.0, 3.0, 4.0], 2, 2).unwrap();
1522        let m2 = Tensor::new_2d(vec![2.0, 3.0, 4.0, 5.0], 2, 2).unwrap();
1523        let result = block_on(elementwise_mul(&Value::Tensor(m1), &Value::Tensor(m2))).unwrap();
1524
1525        if let Value::Tensor(m) = result {
1526            assert_eq!(m.materialize_f64(), vec![2.0, 6.0, 12.0, 20.0]);
1527        } else {
1528            panic!("Expected matrix result");
1529        }
1530    }
1531
1532    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1533    #[test]
1534    fn test_elementwise_div_with_zero() {
1535        let result = block_on(elementwise_div(&Value::Num(5.0), &Value::Num(0.0))).unwrap();
1536        if let Value::Num(n) = result {
1537            assert!(n.is_infinite() && n.is_sign_positive());
1538        } else {
1539            panic!("Expected numeric result");
1540        }
1541    }
1542
1543    #[test]
1544    fn elementwise_div_preserves_native_single_and_nd_shape() {
1545        let shape = vec![1, 2, 2];
1546        let tensor = Tensor::from_f32(vec![1.0, -2.0, 0.0, 4.0], shape.clone()).unwrap();
1547        let Value::Tensor(result) =
1548            block_on(elementwise_div(&Value::Tensor(tensor), &Value::Num(2.0))).expect("div")
1549        else {
1550            panic!("expected tensor");
1551        };
1552        assert_eq!(result.shape, shape);
1553        assert_eq!(
1554            result.into_numeric_storage().expect("storage"),
1555            NumericStorage::F32(vec![0.5, -1.0, 0.0, 2.0])
1556        );
1557    }
1558
1559    #[test]
1560    fn elementwise_div_mixed_floating_tensors_returns_single() {
1561        let single = Tensor::from_f32(vec![2.0, -8.0], vec![1, 2]).unwrap();
1562        let double = Tensor::new(vec![4.0, 2.0], vec![1, 2]).unwrap();
1563
1564        let Value::Tensor(left_single) = block_on(elementwise_div(
1565            &Value::Tensor(single.clone()),
1566            &Value::Tensor(double.clone()),
1567        ))
1568        .expect("div") else {
1569            panic!("expected tensor");
1570        };
1571        assert_eq!(
1572            left_single.into_numeric_storage().expect("storage"),
1573            NumericStorage::F32(vec![0.5, -4.0])
1574        );
1575
1576        let Value::Tensor(right_single) = block_on(elementwise_div(
1577            &Value::Tensor(double),
1578            &Value::Tensor(single),
1579        ))
1580        .expect("div") else {
1581            panic!("expected tensor");
1582        };
1583        assert_eq!(
1584            right_single.into_numeric_storage().expect("storage"),
1585            NumericStorage::F32(vec![2.0, -0.25])
1586        );
1587    }
1588
1589    #[test]
1590    fn elementwise_div_native_single_preserves_zero_policy() {
1591        let tensor = Tensor::from_f32(vec![2.0, -2.0, 0.0], vec![1, 3]).unwrap();
1592        let Value::Tensor(result) =
1593            block_on(elementwise_div(&Value::Tensor(tensor), &Value::Num(0.0))).expect("div")
1594        else {
1595            panic!("expected tensor");
1596        };
1597        let NumericStorage::F32(values) = result.into_numeric_storage().expect("storage") else {
1598            panic!("expected single storage");
1599        };
1600        assert!(values[0].is_infinite() && values[0].is_sign_positive());
1601        assert!(values[1].is_infinite() && values[1].is_sign_negative());
1602        assert!(values[2].is_infinite() && values[2].is_sign_positive());
1603    }
1604
1605    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1606    #[test]
1607    fn test_elementwise_pow() {
1608        let matrix = Tensor::new_2d(vec![2.0, 3.0, 4.0, 5.0], 2, 2).unwrap();
1609        let result = elementwise_pow(&Value::Tensor(matrix), &Value::Num(2.0)).unwrap();
1610
1611        if let Value::Tensor(m) = result {
1612            assert_eq!(m.materialize_f64(), vec![4.0, 9.0, 16.0, 25.0]);
1613        } else {
1614            panic!("Expected matrix result");
1615        }
1616    }
1617
1618    #[test]
1619    fn elementwise_pow_preserves_native_single_and_nd_shape() {
1620        let shape = vec![1, 2, 2];
1621        let tensor = Tensor::from_f32(vec![2.0, 3.0, 4.0, 5.0], shape.clone()).unwrap();
1622        let Value::Tensor(result) =
1623            elementwise_pow(&Value::Tensor(tensor), &Value::Num(2.0)).expect("power")
1624        else {
1625            panic!("expected tensor");
1626        };
1627        assert_eq!(result.shape, shape);
1628        assert_eq!(
1629            result.into_numeric_storage().expect("storage"),
1630            NumericStorage::F32(vec![4.0, 9.0, 16.0, 25.0])
1631        );
1632    }
1633
1634    #[test]
1635    fn elementwise_pow_mixed_floating_tensors_returns_single() {
1636        let single = Tensor::from_f32(vec![2.0, 4.0], vec![1, 2]).unwrap();
1637        let double = Tensor::new(vec![3.0, 0.5], vec![1, 2]).unwrap();
1638
1639        let Value::Tensor(single_base) = elementwise_pow(
1640            &Value::Tensor(single.clone()),
1641            &Value::Tensor(double.clone()),
1642        )
1643        .expect("power") else {
1644            panic!("expected tensor");
1645        };
1646        assert_eq!(
1647            single_base.into_numeric_storage().expect("storage"),
1648            NumericStorage::F32(vec![8.0, 2.0])
1649        );
1650
1651        let Value::Tensor(single_exponent) =
1652            elementwise_pow(&Value::Tensor(double), &Value::Tensor(single)).expect("power")
1653        else {
1654            panic!("expected tensor");
1655        };
1656        assert_eq!(
1657            single_exponent.into_numeric_storage().expect("storage"),
1658            NumericStorage::F32(vec![9.0, 0.0625])
1659        );
1660    }
1661
1662    #[test]
1663    fn elementwise_pow_double_scalar_base_with_single_exponent_returns_single() {
1664        let exponent = Tensor::from_f32(vec![1.0, 2.0, 3.0], vec![1, 3]).unwrap();
1665        let Value::Tensor(result) =
1666            elementwise_pow(&Value::Num(2.0), &Value::Tensor(exponent)).expect("power")
1667        else {
1668            panic!("expected tensor");
1669        };
1670        assert_eq!(
1671            result.into_numeric_storage().expect("storage"),
1672            NumericStorage::F32(vec![2.0, 4.0, 8.0])
1673        );
1674    }
1675
1676    #[test]
1677    fn elementwise_neg_preserves_all_typed_integer_classes_and_shape() {
1678        let cases = [
1679            (
1680                IntegerStorage::I8(vec![i8::MIN, -2, 0, i8::MAX]),
1681                IntegerStorage::I8(vec![i8::MAX, 2, 0, -i8::MAX]),
1682            ),
1683            (
1684                IntegerStorage::I16(vec![i16::MIN, -2, 0, i16::MAX]),
1685                IntegerStorage::I16(vec![i16::MAX, 2, 0, -i16::MAX]),
1686            ),
1687            (
1688                IntegerStorage::I32(vec![i32::MIN, -2, 0, i32::MAX]),
1689                IntegerStorage::I32(vec![i32::MAX, 2, 0, -i32::MAX]),
1690            ),
1691            (
1692                IntegerStorage::I64(vec![i64::MIN, -2, 0, i64::MAX]),
1693                IntegerStorage::I64(vec![i64::MAX, 2, 0, -i64::MAX]),
1694            ),
1695            (
1696                IntegerStorage::U8(vec![0, 2, u8::MAX]),
1697                IntegerStorage::U8(vec![0, 0, 0]),
1698            ),
1699            (
1700                IntegerStorage::U16(vec![0, 2, u16::MAX]),
1701                IntegerStorage::U16(vec![0, 0, 0]),
1702            ),
1703            (
1704                IntegerStorage::U32(vec![0, 2, u32::MAX]),
1705                IntegerStorage::U32(vec![0, 0, 0]),
1706            ),
1707            (
1708                IntegerStorage::U64(vec![0, 2, u64::MAX]),
1709                IntegerStorage::U64(vec![0, 0, 0]),
1710            ),
1711        ];
1712        for (input, expected) in cases {
1713            let shape = vec![1, expected.len(), 1];
1714            let tensor = Tensor::new_integer(input, shape.clone()).expect("integer tensor");
1715            let Value::Tensor(result) = elementwise_neg(&Value::Tensor(tensor)).expect("neg")
1716            else {
1717                panic!("expected tensor");
1718            };
1719            assert_eq!(result.shape, shape);
1720            assert_eq!(result.integer_storage(), Some(&expected));
1721        }
1722    }
1723
1724    #[test]
1725    fn elementwise_neg_preserves_scalar_integer_class() {
1726        assert_eq!(
1727            elementwise_neg(&Value::Int(IntValue::I64(i64::MIN))).expect("neg"),
1728            Value::Int(IntValue::I64(i64::MAX))
1729        );
1730        assert_eq!(
1731            elementwise_neg(&Value::Int(IntValue::U64(u64::MAX))).expect("neg"),
1732            Value::Int(IntValue::U64(0))
1733        );
1734    }
1735
1736    #[test]
1737    fn elementwise_neg_preserves_native_single_storage_and_shape() {
1738        let shape = vec![1, 2, 2];
1739        let tensor =
1740            Tensor::from_f32(vec![1.25, -2.5, f32::INFINITY, -0.0], shape.clone()).unwrap();
1741        let Value::Tensor(result) = elementwise_neg(&Value::Tensor(tensor)).expect("neg") else {
1742            panic!("expected tensor");
1743        };
1744        assert_eq!(result.shape, shape);
1745        assert_eq!(
1746            result.into_numeric_storage().expect("storage"),
1747            NumericStorage::F32(vec![-1.25, 2.5, f32::NEG_INFINITY, 0.0])
1748        );
1749    }
1750
1751    #[test]
1752    fn transitional_elementwise_helpers_preserve_exact_integer_storage() {
1753        let lhs = Tensor::new_integer(
1754            IntegerStorage::U64(vec![u64::MAX, (1_u64 << 63) + 1]),
1755            vec![1, 2],
1756        )
1757        .expect("lhs");
1758        let rhs = Tensor::new_integer(IntegerStorage::U64(vec![1, 2]), vec![1, 2]).expect("rhs");
1759
1760        let Value::Tensor(product) = block_on(elementwise_mul(
1761            &Value::Tensor(lhs.clone()),
1762            &Value::Tensor(rhs),
1763        ))
1764        .expect("mul") else {
1765            panic!("expected integer tensor product");
1766        };
1767        assert_eq!(
1768            product.integer_storage(),
1769            Some(&IntegerStorage::U64(vec![u64::MAX, u64::MAX]))
1770        );
1771
1772        let Value::Tensor(quotient) =
1773            block_on(elementwise_div(&Value::Tensor(lhs), &Value::Num(2.0))).expect("div")
1774        else {
1775            panic!("expected integer tensor quotient");
1776        };
1777        assert_eq!(
1778            quotient.integer_storage(),
1779            Some(&IntegerStorage::U64(vec![1_u64 << 63, (1_u64 << 62) + 1]))
1780        );
1781    }
1782
1783    #[test]
1784    fn elementwise_integer_operations_read_typed_storage_not_poisoned_mirrors() {
1785        let input = Tensor::new_integer(IntegerStorage::I64(vec![2, 3]), vec![1, 2])
1786            .expect("integer tensor");
1787
1788        let Value::Tensor(product) = block_on(elementwise_mul(
1789            &Value::Tensor(input.clone()),
1790            &Value::Num(0.5),
1791        ))
1792        .expect("product") else {
1793            panic!("expected tensor");
1794        };
1795        assert_eq!(
1796            product.integer_storage(),
1797            Some(&IntegerStorage::I64(vec![1, 2]))
1798        );
1799
1800        let Value::Tensor(quotient) = block_on(elementwise_div(
1801            &Value::Num(6.0),
1802            &Value::Tensor(input.clone()),
1803        ))
1804        .expect("quotient") else {
1805            panic!("expected tensor");
1806        };
1807        assert_eq!(
1808            quotient.integer_storage(),
1809            Some(&IntegerStorage::I64(vec![3, 2]))
1810        );
1811
1812        let error = elementwise_pow(&Value::Tensor(input), &Value::Num(0.5))
1813            .expect_err("fractional integer exponent must reject");
1814        assert!(error.contains("nonnegative integer values"));
1815    }
1816
1817    #[test]
1818    fn transitional_power_helpers_preserve_exact_scalar_and_array_integers() {
1819        let scalar_power =
1820            power(&Value::Int(IntValue::U64(u64::MAX)), &Value::Num(1.0)).expect("scalar power");
1821        assert_eq!(scalar_power, Value::Int(IntValue::U64(u64::MAX)));
1822
1823        let scalar_tensor = Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX]), vec![1, 1])
1824            .expect("scalar tensor");
1825        let scalar_tensor_power =
1826            power(&Value::Tensor(scalar_tensor), &Value::Num(1.0)).expect("tensor scalar power");
1827        assert_eq!(scalar_tensor_power, Value::Int(IntValue::U64(u64::MAX)));
1828
1829        let complex_base =
1830            Tensor::new_integer(IntegerStorage::U8(vec![3]), vec![1, 1]).expect("complex base");
1831        let complex_power = power(&Value::Tensor(complex_base), &Value::Complex(1.0, 0.0))
1832            .expect("complex exponent power");
1833        let Value::Complex(re, im) = complex_power else {
1834            panic!("expected complex scalar");
1835        };
1836        assert!((re - 3.0).abs() < 1e-12);
1837        assert_eq!(im, 0.0);
1838
1839        let base =
1840            Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX, 2]), vec![1, 2]).expect("base");
1841        let exponent =
1842            Tensor::new_integer(IntegerStorage::U64(vec![1, 64]), vec![1, 2]).expect("exponent");
1843        let Value::Tensor(result) =
1844            elementwise_pow(&Value::Tensor(base), &Value::Tensor(exponent)).expect("pow")
1845        else {
1846            panic!("expected integer tensor power");
1847        };
1848        assert_eq!(
1849            result.integer_storage(),
1850            Some(&IntegerStorage::U64(vec![u64::MAX, u64::MAX]))
1851        );
1852    }
1853
1854    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1855    #[test]
1856    fn test_dimension_mismatch() {
1857        let m1 = Tensor::new_2d(vec![1.0, 2.0], 1, 2).unwrap();
1858        let m2 = Tensor::new_2d(vec![1.0, 2.0, 3.0, 4.0], 2, 2).unwrap();
1859
1860        assert!(block_on(elementwise_mul(&Value::Tensor(m1), &Value::Tensor(m2))).is_err());
1861    }
1862}