Skip to main content

runmat_runtime/builtins/logical/
ops.rs

1//! MATLAB-compatible `logical` builtin with GPU-aware semantics for RunMat.
2
3use log::trace;
4use runmat_accelerate_api::{self, AccelProvider, GpuTensorHandle, HostTensorView};
5use runmat_builtins::{
6    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
7    BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
8    BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
9    BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
10    BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
11    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
12    ResolveContext, Type,
13};
14use runmat_macros::runtime_builtin;
15#[cfg(test)]
16use runmat_value::ComplexTensor;
17use runmat_value::{CharArray, LogicalArray, StringArray, Tensor, Value};
18
19use crate::builtins::common::{
20    gpu_helpers,
21    shape::{canonical_scalar_shape, normalize_scalar_shape},
22    spec::{
23        BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
24        ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
25    },
26    tensor,
27};
28use crate::builtins::logical::type_resolvers::logical_like;
29
30use crate::{build_runtime_error, BuiltinResult, RuntimeError};
31
32#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::logical::ops")]
33pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
34    name: "logical",
35    op_kind: GpuOpKind::Elementwise,
36    supported_precisions: &[ScalarType::F32, ScalarType::F64],
37    broadcast: BroadcastSemantics::Matlab,
38    provider_hooks: &[ProviderHook::Binary {
39        name: "elem_ne",
40        commutative: true,
41    }],
42    constant_strategy: ConstantStrategy::InlineLiteral,
43    residency: ResidencyPolicy::NewHandle,
44    nan_mode: ReductionNaN::Include,
45    two_pass_threshold: None,
46    workgroup_size: None,
47    accepts_nan_mode: false,
48    notes: "Preferred path issues elem_ne(X, 0) on the device; missing hooks trigger a gather → host cast → re-upload sequence flagged as logical.",
49};
50
51#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::logical::ops")]
52pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
53    name: "logical",
54    shape: ShapeRequirements::BroadcastCompatible,
55    constant_strategy: ConstantStrategy::InlineLiteral,
56    elementwise: None,
57    reduction: None,
58    emits_nan: false,
59    notes: "Fusion support will arrive alongside a dedicated WGSL template; today the builtin executes outside fusion plans.",
60};
61
62const BUILTIN_NAME: &str = "logical";
63
64const LOGICAL_STRING_ARRAY_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
65    id: "logical-string-array-input",
66    mode: BuiltinExtensionMode::RunMatOnly,
67    description: "logical with string-array input is a RunMat extension",
68    error_identifier: Some("RunMat:compatibility:LogicalStringArrayInputExtension"),
69};
70const LOGICAL_SYMBOLIC_CONSTANT_EXTENSION: BuiltinExtensionDescriptor =
71    BuiltinExtensionDescriptor {
72        id: "logical-symbolic-constant-input",
73        mode: BuiltinExtensionMode::RunMatOnly,
74        description: "logical with a symbolic numeric constant is a RunMat extension",
75        error_identifier: Some("RunMat:compatibility:LogicalSymbolicConstantInputExtension"),
76    };
77pub const LOGICAL_EXTENSIONS: [BuiltinExtensionDescriptor; 2] = [
78    LOGICAL_STRING_ARRAY_EXTENSION,
79    LOGICAL_SYMBOLIC_CONSTANT_EXTENSION,
80];
81
82const LOGICAL_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 1] =
83    [BuiltinIntegerInputCapability {
84        name: "A",
85        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
86        availability: BuiltinIntegerInputAvailability::Documented,
87        scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
88        notes: "All eight real integer classes convert elementwise without floating materialization; zero becomes false and every nonzero value becomes true.",
89    }];
90pub const LOGICAL_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
91    [BuiltinIntegerCapabilityDescriptor {
92        form: "tf = logical(integer_A)",
93        inputs: &LOGICAL_INTEGER_INPUTS,
94        computation_domain: BuiltinIntegerComputationDomain::Predicate,
95        output_class: BuiltinIntegerOutputClassRule::Logical,
96        overflow: BuiltinIntegerOverflowRule::NotApplicable,
97        backend: BuiltinIntegerBackendRule::HostAndGpu,
98        overload: BuiltinIntegerOverloadKind::ElementwiseShapePreserving,
99        notes: "Host conversion reads authoritative integer storage exactly; resident conversion uses a validated owning-provider path or exact gather and class-preserving restoration.",
100    }];
101
102const LOGICAL_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
103    name: "tf",
104    ty: BuiltinParamType::LogicalArray,
105    arity: BuiltinParamArity::Required,
106    default: None,
107    description: "Logical-converted result.",
108}];
109
110const LOGICAL_INPUTS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
111    name: "A",
112    ty: BuiltinParamType::Any,
113    arity: BuiltinParamArity::Required,
114    default: None,
115    description: "Input value to convert.",
116}];
117
118const LOGICAL_SIGNATURES: [BuiltinSignatureDescriptor; 1] = [BuiltinSignatureDescriptor {
119    label: "tf = logical(A)",
120    inputs: &LOGICAL_INPUTS,
121    outputs: &LOGICAL_OUTPUT,
122}];
123
124const LOGICAL_ERROR_TOO_MANY_INPUTS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
125    code: "RM.LOGICAL.TOO_MANY_INPUTS",
126    identifier: Some("RunMat:logical:TooManyInputs"),
127    when: "More than one input argument is provided.",
128    message: "logical: too many input arguments",
129};
130
131const LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
132    code: "RM.LOGICAL.CONVERSION_NOT_POSSIBLE",
133    identifier: Some("RunMat:logical:ConversionNotPossible"),
134    when: "Input type cannot be converted to logical.",
135    message: "logical: conversion to logical is not possible for this input type",
136};
137
138const LOGICAL_ERROR_GPU_GATHER_FAILED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
139    code: "RM.LOGICAL.GPU_GATHER_FAILED",
140    identifier: Some("RunMat:logical:GpuGatherFailed"),
141    when: "GPU input gather fails during host fallback.",
142    message: "logical: failed to gather gpuArray input",
143};
144
145const LOGICAL_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
146    code: "RM.LOGICAL.INTERNAL",
147    identifier: Some("RunMat:logical:InternalError"),
148    when: "Internal logical buffer materialization fails.",
149    message: "logical: internal conversion error",
150};
151
152const LOGICAL_ERRORS: [BuiltinErrorDescriptor; 4] = [
153    LOGICAL_ERROR_TOO_MANY_INPUTS,
154    LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE,
155    LOGICAL_ERROR_GPU_GATHER_FAILED,
156    LOGICAL_ERROR_INTERNAL,
157];
158
159pub const LOGICAL_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
160    signatures: &LOGICAL_SIGNATURES,
161    output_mode: BuiltinOutputMode::Fixed,
162    completion_policy: BuiltinCompletionPolicy::Public,
163    errors: &LOGICAL_ERRORS,
164};
165
166fn logical_type(args: &[Type], _context: &ResolveContext) -> Type {
167    args.first().map(logical_like).unwrap_or(Type::logical())
168}
169
170fn logical_error_with_message(
171    message: impl Into<String>,
172    error: &'static BuiltinErrorDescriptor,
173) -> RuntimeError {
174    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
175    if let Some(identifier) = error.identifier {
176        builder = builder.with_identifier(identifier);
177    }
178    builder.build()
179}
180
181#[runtime_builtin(
182    name = "logical",
183    category = "logical",
184    summary = "Convert scalars, arrays, and gpuArray values to logical outputs.",
185    keywords = "logical,boolean,gpuArray,mask,conversion",
186    accel = "unary",
187    type_resolver(logical_type),
188    descriptor(crate::builtins::logical::ops::LOGICAL_DESCRIPTOR),
189    extensions(crate::builtins::logical::ops::LOGICAL_EXTENSIONS),
190    integer_capabilities(crate::builtins::logical::ops::LOGICAL_INTEGER_CAPABILITIES),
191    builtin_path = "crate::builtins::logical::ops"
192)]
193async fn logical_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
194    if !rest.is_empty() {
195        return Err(logical_error_with_message(
196            LOGICAL_ERROR_TOO_MANY_INPUTS.message,
197            &LOGICAL_ERROR_TOO_MANY_INPUTS,
198        ));
199    }
200    convert_value_to_logical(value).await
201}
202
203async fn convert_value_to_logical(value: Value) -> BuiltinResult<Value> {
204    match value {
205        Value::Bool(_) | Value::LogicalArray(_) => Ok(value),
206        Value::Num(n) if n.is_nan() => Err(conversion_error("NaN")),
207        Value::Num(n) => Ok(Value::Bool(n != 0.0)),
208        Value::Int(i) => Ok(Value::Bool(!i.is_zero())),
209        Value::Complex(_, _) => Err(conversion_error("complex")),
210        Value::Tensor(tensor) => logical_from_tensor(tensor),
211        Value::SparseTensor(sparse) => logical_from_sparse_tensor(sparse),
212        Value::ComplexTensor(_) => Err(conversion_error("complex")),
213        Value::CharArray(chars) => logical_from_char_array(chars),
214        Value::StringArray(strings) => {
215            crate::compatibility::ensure_builtin_extension_enabled(
216                &LOGICAL_STRING_ARRAY_EXTENSION,
217                BUILTIN_NAME,
218            )?;
219            logical_from_string_array(strings)
220        }
221        Value::GpuTensor(handle) => logical_from_gpu(handle).await,
222        Value::String(_) => Err(conversion_error("string")),
223        Value::Symbolic(expr) => expr
224            .numeric_constant_value()
225            .map(|value| Value::Bool(value != 0.0))
226            .ok_or_else(|| conversion_error("sym")),
227        Value::SymbolicArray(_) => Err(conversion_error("sym")),
228        Value::Cell(_) => Err(conversion_error("cell")),
229        Value::Struct(_) => Err(conversion_error("struct")),
230        Value::ObjectArray(array) => Err(conversion_error(array.class_name())),
231        Value::Object(obj) => Err(conversion_error(&obj.class_name)),
232        Value::HandleObject(handle) => Err(conversion_error(&handle.class_name)),
233        Value::Listener(_) => Err(conversion_error("event.listener")),
234        Value::FunctionHandle(_)
235        | Value::ExternalFunctionHandle(_)
236        | Value::MethodFunctionHandle(_)
237        | Value::BoundFunctionHandle { .. }
238        | Value::Closure(_) => Err(conversion_error("function_handle")),
239        Value::ClassRef(_) => Err(conversion_error("meta.class")),
240        Value::MException(_)
241        | Value::Future(_)
242        | Value::Task(_)
243        | Value::Pool(_)
244        | Value::Job(_) => Err(conversion_error("MException")),
245        Value::Foreign(_) => Err(conversion_error("foreign")),
246        Value::OutputList(_) => Err(conversion_error("OutputList")),
247    }
248}
249
250fn logical_from_tensor(tensor: Tensor) -> BuiltinResult<Value> {
251    if tensor.integer_storage().is_none()
252        && tensor.materialize_f64().iter().any(|value| value.is_nan())
253    {
254        return Err(conversion_error("NaN"));
255    }
256    let buffer = LogicalBuffer::from_real_tensor(&tensor);
257    logical_buffer_to_host(buffer)
258}
259
260fn logical_from_sparse_tensor(sparse: runmat_value::SparseTensor) -> BuiltinResult<Value> {
261    if sparse.is_logical() {
262        return Ok(Value::SparseTensor(sparse));
263    }
264    if sparse.integer_storage().is_none()
265        && (0..sparse.nnz()).any(|index| {
266            sparse
267                .numeric_value_at(index)
268                .is_some_and(|value| match value {
269                    runmat_value::NumericScalar::F64(value) => value.is_nan(),
270                    runmat_value::NumericScalar::F32(value) => value.is_nan(),
271                    _ => false,
272                })
273        })
274    {
275        return Err(conversion_error("NaN"));
276    }
277    let mut col_ptrs = Vec::with_capacity(sparse.cols.saturating_add(1));
278    let mut row_indices = Vec::new();
279    col_ptrs.push(0);
280    for col in 0..sparse.cols {
281        for index in sparse.col_ptrs[col]..sparse.col_ptrs[col + 1] {
282            if !sparse
283                .numeric_value_at(index)
284                .expect("validated sparse storage index")
285                .is_zero()
286            {
287                row_indices.push(sparse.row_indices[index]);
288            }
289        }
290        col_ptrs.push(row_indices.len());
291    }
292    runmat_value::SparseTensor::new_logical(sparse.rows, sparse.cols, col_ptrs, row_indices)
293        .map(Value::SparseTensor)
294        .map_err(|err| {
295            logical_error_with_message(
296                format!("logical: failed to convert sparse input: {err}"),
297                &LOGICAL_ERROR_INTERNAL,
298            )
299        })
300}
301
302fn logical_from_char_array(chars: CharArray) -> BuiltinResult<Value> {
303    let buffer = LogicalBuffer::from_char_array(&chars);
304    logical_buffer_to_host(buffer)
305}
306
307fn logical_from_string_array(strings: StringArray) -> BuiltinResult<Value> {
308    let bits: Vec<u8> = strings
309        .data
310        .iter()
311        .map(|s| if s.is_empty() { 0 } else { 1 })
312        .collect();
313    let shape = canonical_shape(&strings.shape, bits.len());
314    logical_buffer_to_host(LogicalBuffer { bits, shape })
315}
316
317async fn logical_from_gpu(handle: GpuTensorHandle) -> BuiltinResult<Value> {
318    if runmat_accelerate_api::handle_is_logical(&handle) {
319        return Ok(Value::GpuTensor(handle));
320    }
321
322    if runmat_accelerate_api::handle_storage(&handle)
323        == runmat_accelerate_api::GpuTensorStorage::ComplexInterleaved
324    {
325        return Err(conversion_error("complex"));
326    }
327    let provider = gpu_helpers::exact_provider_for_handle(&handle);
328
329    if let Some(p) = provider {
330        let contains_nan = match provider_input_contains_nan(p, &handle).await {
331            Ok(contains_nan) => contains_nan,
332            Err(_) => {
333                let host =
334                    gpu_helpers::download_value_preserving_residency_async(p, &handle).await?;
335                value_contains_nan(&host)
336            }
337        };
338        if contains_nan {
339            return Err(conversion_error("NaN"));
340        }
341        match p.logical_islogical(&handle) {
342            Ok(true) => {
343                runmat_accelerate_api::set_handle_logical(&handle, true);
344                return Ok(Value::GpuTensor(handle));
345            }
346            Ok(false) => {}
347            Err(err) => {
348                trace!("logical: provider logical_islogical hook unavailable, falling back ({err})")
349            }
350        }
351        if let Some(mut result) = try_gpu_cast(p, &handle).await {
352            copy_logical_provenance(&mut result, &handle);
353            return Ok(gpu_helpers::logical_gpu_value(result));
354        } else {
355            trace!(
356                "logical: provider elem_ne/zeros_like unavailable for buffer {} – gathering",
357                handle.buffer_id
358            );
359        }
360    }
361
362    let tensor = gpu_helpers::gather_tensor_async(&handle)
363        .await
364        .map_err(|err| {
365            logical_error_with_message(
366                format!("{BUILTIN_NAME}: {err}"),
367                &LOGICAL_ERROR_GPU_GATHER_FAILED,
368            )
369        })?;
370    let buffer = LogicalBuffer::from_real_tensor(&tensor);
371    logical_buffer_to_gpu(
372        buffer,
373        provider,
374        runmat_accelerate_api::handle_is_explicit(&handle),
375    )
376}
377
378async fn provider_input_contains_nan(
379    provider: &'static dyn runmat_accelerate_api::AccelProvider,
380    source: &GpuTensorHandle,
381) -> BuiltinResult<bool> {
382    let mask = provider
383        .logical_isnan(source)
384        .map_err(|error| logical_error_with_message(error.to_string(), &LOGICAL_ERROR_INTERNAL))?;
385    let mask_valid = !gpu_helpers::same_gpu_handle(&mask, source)
386        && mask.shape == source.shape
387        && mask.device_id == source.device_id
388        && gpu_helpers::exact_provider_for_handle(&mask)
389            .is_some_and(|owner| std::ptr::eq(owner, provider));
390    if !mask_valid {
391        gpu_helpers::free_unprotected_exact_owner(&mask, &[source]);
392        return Err(logical_error_with_message(
393            "logical: provider returned a malformed NaN mask",
394            &LOGICAL_ERROR_INTERNAL,
395        ));
396    }
397    let maximum = provider.reduce_max(&mask).await.map_err(|error| {
398        gpu_helpers::free_unprotected_exact_owner(&mask, &[source]);
399        logical_error_with_message(error.to_string(), &LOGICAL_ERROR_INTERNAL)
400    })?;
401    let maximum_valid = !gpu_helpers::same_gpu_handle(&maximum, source)
402        && !gpu_helpers::same_gpu_handle(&maximum, &mask)
403        && maximum.shape.iter().product::<usize>() == 1
404        && maximum.device_id == source.device_id
405        && gpu_helpers::exact_provider_for_handle(&maximum)
406            .is_some_and(|owner| std::ptr::eq(owner, provider));
407    if !maximum_valid {
408        gpu_helpers::free_unprotected_exact_owner(&maximum, &[source, &mask]);
409        gpu_helpers::free_unprotected_exact_owner(&mask, &[source]);
410        return Err(logical_error_with_message(
411            "logical: provider returned a malformed NaN reduction",
412            &LOGICAL_ERROR_INTERNAL,
413        ));
414    }
415    let downloaded = provider.download(&maximum).await.map_err(|error| {
416        gpu_helpers::free_unprotected_exact_owner(&maximum, &[source, &mask]);
417        gpu_helpers::free_unprotected_exact_owner(&mask, &[source]);
418        logical_error_with_message(error.to_string(), &LOGICAL_ERROR_INTERNAL)
419    })?;
420    gpu_helpers::free_unprotected_exact_owner(&maximum, &[source, &mask]);
421    gpu_helpers::free_unprotected_exact_owner(&mask, &[source]);
422    Ok(downloaded.data.first().is_some_and(|value| *value != 0.0))
423}
424
425fn value_contains_nan(value: &Value) -> bool {
426    match value {
427        Value::Num(value) => value.is_nan(),
428        Value::Tensor(tensor) => {
429            tensor.integer_storage().is_none()
430                && tensor.materialize_f64().iter().any(|value| value.is_nan())
431        }
432        _ => false,
433    }
434}
435
436fn logical_buffer_to_host(buffer: LogicalBuffer) -> BuiltinResult<Value> {
437    let LogicalBuffer { bits, shape } = buffer;
438    if tensor::element_count(&shape) == 1 && bits.len() == 1 {
439        Ok(Value::Bool(bits[0] != 0))
440    } else {
441        LogicalArray::new(bits, shape)
442            .map(Value::LogicalArray)
443            .map_err(|e| {
444                logical_error_with_message(format!("logical: {e}"), &LOGICAL_ERROR_INTERNAL)
445            })
446    }
447}
448
449fn logical_buffer_to_gpu(
450    buffer: LogicalBuffer,
451    provider: Option<&'static dyn AccelProvider>,
452    explicit: bool,
453) -> BuiltinResult<Value> {
454    if let Some(p) = provider {
455        let floats: Vec<f64> = buffer
456            .bits
457            .iter()
458            .map(|&b| if b != 0 { 1.0 } else { 0.0 })
459            .collect();
460        let view = HostTensorView {
461            data: &floats,
462            shape: &buffer.shape,
463        };
464        match p.upload(&view) {
465            Ok(mut handle) => {
466                if explicit {
467                    runmat_accelerate_api::mark_handle_explicit(&mut handle);
468                } else {
469                    runmat_accelerate_api::mark_handle_automatic(&mut handle);
470                }
471                Ok(gpu_helpers::logical_gpu_value(handle))
472            }
473            Err(err) => {
474                trace!("logical: upload failed during fallback path ({err})");
475                if explicit {
476                    Err(logical_error_with_message(
477                        format!("logical: failed to preserve explicit gpuArray residency: {err}"),
478                        &LOGICAL_ERROR_INTERNAL,
479                    ))
480                } else {
481                    logical_buffer_to_host(buffer)
482                }
483            }
484        }
485    } else if explicit {
486        Err(logical_error_with_message(
487            "logical: no exact owner for explicit gpuArray input",
488            &LOGICAL_ERROR_GPU_GATHER_FAILED,
489        ))
490    } else {
491        logical_buffer_to_host(buffer)
492    }
493}
494
495async fn try_gpu_cast(
496    provider: &'static dyn AccelProvider,
497    input: &GpuTensorHandle,
498) -> Option<GpuTensorHandle> {
499    let zeros = provider.zeros_like(input).ok()?;
500    let zeros_valid = zeros.shape == input.shape
501        && zeros.device_id == input.device_id
502        && !gpu_helpers::same_gpu_handle(&zeros, input)
503        && runmat_accelerate_api::handle_storage(&zeros)
504            == runmat_accelerate_api::GpuTensorStorage::Real
505        && runmat_accelerate_api::handle_precision(&zeros)
506            == runmat_accelerate_api::handle_precision(input)
507        && runmat_accelerate_api::handle_integer_type(&zeros)
508            == runmat_accelerate_api::handle_integer_type(input)
509        && runmat_accelerate_api::handle_is_logical(&zeros)
510            == runmat_accelerate_api::handle_is_logical(input)
511        && gpu_helpers::exact_provider_for_handle(&zeros)
512            .is_some_and(|owner| std::ptr::eq(owner, provider));
513    if !zeros_valid {
514        gpu_helpers::free_unprotected_exact_owner(&zeros, &[input]);
515        return None;
516    }
517    let result = provider
518        .elem_ne(input, &zeros)
519        .await
520        .ok()
521        .and_then(|output| {
522            if valid_logical_gpu_output(&output, input, provider) {
523                Some(output)
524            } else {
525                gpu_helpers::free_unprotected_exact_owner(&output, &[input, &zeros]);
526                None
527            }
528        });
529    let _ = provider.free(&zeros);
530    result
531}
532
533fn copy_logical_provenance(output: &mut GpuTensorHandle, input: &GpuTensorHandle) {
534    runmat_accelerate_api::set_handle_provenance(
535        output,
536        runmat_accelerate_api::handle_provenance(input)
537            .unwrap_or(runmat_accelerate_api::GpuHandleProvenance::Automatic),
538    );
539}
540
541fn valid_logical_gpu_output(
542    output: &GpuTensorHandle,
543    input: &GpuTensorHandle,
544    provider: &'static dyn AccelProvider,
545) -> bool {
546    output.shape == input.shape
547        && output.device_id == input.device_id
548        && !gpu_helpers::same_gpu_handle(output, input)
549        && runmat_accelerate_api::handle_storage(output)
550            == runmat_accelerate_api::GpuTensorStorage::Real
551        && runmat_accelerate_api::handle_integer_type(output).is_none()
552        && gpu_helpers::exact_provider_for_handle(output)
553            .is_some_and(|owner| std::ptr::eq(owner, provider))
554}
555
556fn conversion_error(type_name: &str) -> RuntimeError {
557    logical_error_with_message(
558        format!(
559            "logical: conversion to logical from {} is not possible",
560            type_name
561        ),
562        &LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE,
563    )
564}
565
566#[derive(Clone)]
567struct LogicalBuffer {
568    bits: Vec<u8>,
569    shape: Vec<usize>,
570}
571
572impl LogicalBuffer {
573    fn from_real_tensor(tensor: &Tensor) -> Self {
574        let bits: Vec<u8> = (0..tensor.len())
575            .map(|index| {
576                u8::from(
577                    !tensor
578                        .numeric_value_at(index)
579                        .expect("tensor storage is structurally valid")
580                        .is_zero(),
581                )
582            })
583            .collect();
584        let shape = canonical_shape(&tensor.shape, bits.len());
585        Self { bits, shape }
586    }
587
588    fn from_char_array(chars: &CharArray) -> Self {
589        let bits: Vec<u8> = chars
590            .data
591            .iter()
592            .map(|&ch| if (ch as u32) != 0 { 1 } else { 0 })
593            .collect();
594        let original_shape = vec![chars.rows, chars.cols];
595        let shape = canonical_shape(&original_shape, bits.len());
596        Self { bits, shape }
597    }
598}
599
600fn canonical_shape(shape: &[usize], len: usize) -> Vec<usize> {
601    if tensor::element_count(shape) == len {
602        return normalize_scalar_shape(shape);
603    }
604    if len == 0 {
605        if shape.len() > 1 {
606            return shape.to_vec();
607        }
608        return vec![0];
609    }
610    if len == 1 {
611        canonical_scalar_shape()
612    } else {
613        vec![len, 1]
614    }
615}
616
617#[cfg(test)]
618pub(crate) mod tests {
619    use super::*;
620    use crate::builtins::common::test_support;
621    use futures::executor::block_on;
622    use runmat_accelerate_api::HostTensorView;
623    use runmat_value::{
624        CellArray, IntValue, IntegerComplexStorage, IntegerStorage, MException, ObjectInstance,
625        SparseTensor, StructValue, SymbolicExpr,
626    };
627
628    fn logical_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
629        block_on(super::logical_builtin(value, rest))
630    }
631
632    fn assert_error_message(err: &crate::RuntimeError, expected: &str) {
633        assert_eq!(err.message(), expected);
634    }
635
636    fn assert_error_contains(err: &crate::RuntimeError, expected: &str) {
637        assert!(
638            err.message().contains(expected),
639            "unexpected error: {}",
640            err.message()
641        );
642    }
643
644    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
645    #[test]
646    fn logical_scalar_num() {
647        let result = logical_builtin(Value::Num(5.0), Vec::new()).expect("logical");
648        assert_eq!(result, Value::Bool(true));
649
650        let zero_result = logical_builtin(Value::Num(0.0), Vec::new()).expect("logical");
651        assert_eq!(zero_result, Value::Bool(false));
652    }
653
654    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
655    #[test]
656    fn logical_converts_symbolic_constants() {
657        let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
658        let nonzero = logical_builtin(Value::Symbolic(SymbolicExpr::constant(2.0)), Vec::new())
659            .expect("logical");
660        assert_eq!(nonzero, Value::Bool(true));
661
662        let zero = logical_builtin(Value::Symbolic(SymbolicExpr::constant(0.0)), Vec::new())
663            .expect("logical");
664        assert_eq!(zero, Value::Bool(false));
665    }
666
667    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
668    #[test]
669    fn logical_rejects_symbolic_variables() {
670        let err = logical_builtin(Value::Symbolic(SymbolicExpr::variable("x")), Vec::new())
671            .expect_err("symbolic variable should not convert");
672
673        assert_eq!(
674            err.identifier(),
675            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
676        );
677        assert!(err.message().contains("logical from sym"));
678    }
679
680    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
681    #[test]
682    fn logical_rejects_nan() {
683        let tensor = Tensor::new(vec![0.0, f64::NAN, -0.0], vec![1, 3]).unwrap();
684        let error = logical_builtin(Value::Tensor(tensor), Vec::new())
685            .expect_err("NaN conversion must fail");
686        assert!(error.message().contains("NaN"));
687    }
688
689    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
690    #[test]
691    fn logical_tensor_matrix() {
692        let tensor = Tensor::new(vec![0.0, 2.0, -3.0, 0.0], vec![2, 2]).unwrap();
693        let result = logical_builtin(Value::Tensor(tensor), Vec::new()).expect("logical");
694        match result {
695            Value::LogicalArray(array) => {
696                assert_eq!(array.shape, vec![2, 2]);
697                assert_eq!(array.data, vec![0, 1, 1, 0]);
698            }
699            other => panic!("expected logical array, got {:?}", other),
700        }
701    }
702
703    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
704    #[test]
705    fn logical_sparse_tensor_preserves_sparse_storage() {
706        let sparse = SparseTensor::new(3, 2, vec![0, 1, 2], vec![1, 2], vec![4.0, -1.0]).unwrap();
707        let result = logical_builtin(Value::SparseTensor(sparse), Vec::new()).expect("logical");
708        match result {
709            Value::SparseTensor(sparse) => {
710                assert!(sparse.is_logical());
711                assert_eq!(sparse.shape(), vec![3, 2]);
712                assert_eq!(sparse.col_ptrs, vec![0, 1, 2]);
713                assert_eq!(sparse.row_indices, vec![1, 2]);
714                assert_eq!(
715                    sparse.to_dense_logical().expect("dense logical").data,
716                    vec![0, 1, 0, 0, 0, 1]
717                );
718            }
719            other => panic!("expected logical sparse tensor, got {other:?}"),
720        }
721    }
722
723    #[test]
724    fn logical_sparse_nan_rejects_like_dense_nan() {
725        let sparse = SparseTensor::new(2, 1, vec![0, 1], vec![0], vec![f64::NAN]).unwrap();
726        let error = logical_builtin(Value::SparseTensor(sparse), Vec::new())
727            .expect_err("sparse NaN must reject");
728        assert_eq!(
729            error.identifier(),
730            Some("RunMat:logical:ConversionNotPossible")
731        );
732    }
733
734    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
735    #[test]
736    fn logical_rejects_complex_conversion() {
737        let complex =
738            ComplexTensor::new(vec![(0.0, 0.0), (1.0, 0.0), (0.0, 2.0)], vec![3, 1]).unwrap();
739        let error = logical_builtin(Value::ComplexTensor(complex), Vec::new())
740            .expect_err("complex conversion must fail");
741        assert!(error.message().contains("complex"));
742    }
743
744    #[test]
745    fn logical_rejects_typed_complex_integer_components() {
746        let storage = IntegerComplexStorage::new(
747            IntegerStorage::U64(vec![0, u64::MAX, 0]),
748            IntegerStorage::U64(vec![0, 0, 1_u64 << 63]),
749        )
750        .expect("matching components");
751        let tensor = ComplexTensor::new_integer(storage, vec![3, 1]).expect("typed complex");
752
753        let error = logical_builtin(Value::ComplexTensor(tensor), Vec::new())
754            .expect_err("complex conversion must fail");
755        assert!(error.message().contains("complex"));
756    }
757
758    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
759    #[test]
760    fn logical_char_array_conversion() {
761        let chars = CharArray::new(vec!['A', '\0', 'C'], 1, 3).unwrap();
762        let result = logical_builtin(Value::CharArray(chars), Vec::new()).expect("logical");
763        match result {
764            Value::LogicalArray(array) => assert_eq!(array.data, vec![1, 0, 1]),
765            other => panic!("expected logical array, got {:?}", other),
766        }
767    }
768
769    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
770    #[test]
771    fn logical_string_error() {
772        let err = logical_builtin(Value::String("runmat".to_string()), Vec::new()).unwrap_err();
773        assert_error_message(
774            &err,
775            "logical: conversion to logical from string is not possible",
776        );
777        assert_eq!(
778            err.identifier(),
779            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
780        );
781    }
782
783    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
784    #[test]
785    fn logical_struct_error() {
786        let mut st = StructValue::new();
787        st.insert("field", Value::Num(1.0));
788        let err = logical_builtin(Value::Struct(st), Vec::new()).unwrap_err();
789        assert_error_contains(&err, "struct");
790        assert_eq!(
791            err.identifier(),
792            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
793        );
794    }
795
796    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
797    #[test]
798    fn logical_cell_error() {
799        let cell = CellArray::new(vec![Value::Num(1.0)], 1, 1).expect("cell creation");
800        let err = logical_builtin(Value::Cell(cell), Vec::new()).unwrap_err();
801        assert_error_message(
802            &err,
803            "logical: conversion to logical from cell is not possible",
804        );
805        assert_eq!(
806            err.identifier(),
807            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
808        );
809    }
810
811    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
812    #[test]
813    fn logical_function_handle_error() {
814        let err = logical_builtin(Value::FunctionHandle("foo".into()), Vec::new()).unwrap_err();
815        assert_error_message(
816            &err,
817            "logical: conversion to logical from function_handle is not possible",
818        );
819        assert_eq!(
820            err.identifier(),
821            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
822        );
823    }
824
825    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
826    #[test]
827    fn logical_object_error() {
828        let obj = ObjectInstance::new("DemoClass".to_string());
829        let err = logical_builtin(Value::Object(obj), Vec::new()).unwrap_err();
830        assert_error_contains(&err, "DemoClass");
831        assert_eq!(
832            err.identifier(),
833            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
834        );
835    }
836
837    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
838    #[test]
839    fn logical_mexception_error() {
840        let mex = MException::new("id:logical".into(), "message".into());
841        let err = logical_builtin(Value::MException(mex), Vec::new()).unwrap_err();
842        assert_error_message(
843            &err,
844            "logical: conversion to logical from MException is not possible",
845        );
846        assert_eq!(
847            err.identifier(),
848            LOGICAL_ERROR_CONVERSION_NOT_POSSIBLE.identifier
849        );
850    }
851
852    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
853    #[test]
854    fn logical_too_many_inputs_error() {
855        let err = logical_builtin(Value::Bool(true), vec![Value::Bool(false)]).unwrap_err();
856        assert_error_message(&err, LOGICAL_ERROR_TOO_MANY_INPUTS.message);
857        assert_eq!(err.identifier(), LOGICAL_ERROR_TOO_MANY_INPUTS.identifier);
858    }
859
860    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
861    #[test]
862    fn logical_gpu_roundtrip() {
863        test_support::with_test_provider(|provider| {
864            let tensor = Tensor::new(vec![0.0, 1.0, -2.0], vec![3, 1]).unwrap();
865            let view = HostTensorView {
866                data: &tensor.materialize_f64(),
867                shape: &tensor.shape,
868            };
869            let handle = provider.upload(&view).expect("upload");
870            let result =
871                logical_builtin(Value::GpuTensor(handle.clone()), Vec::new()).expect("logical");
872            let gathered = test_support::gather(result.clone()).expect("gather");
873            assert_eq!(gathered.materialize_f64(), vec![0.0, 1.0, 1.0]);
874            if let Value::GpuTensor(out) = result {
875                assert!(runmat_accelerate_api::handle_is_logical(&out));
876            } else {
877                panic!("expected gpu tensor output");
878            }
879        });
880    }
881
882    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
883    #[test]
884    fn logical_gpu_passthrough_for_logical_handle() {
885        test_support::with_test_provider(|provider| {
886            let tensor = Tensor::new(vec![0.0, 1.0], vec![2, 1]).unwrap();
887            let view = HostTensorView {
888                data: &tensor.materialize_f64(),
889                shape: &tensor.shape,
890            };
891            let handle = provider.upload(&view).expect("upload");
892            runmat_accelerate_api::set_handle_logical(&handle, true);
893            let result =
894                logical_builtin(Value::GpuTensor(handle.clone()), Vec::new()).expect("logical");
895            match result {
896                Value::GpuTensor(out) => assert_eq!(out, handle),
897                other => panic!("expected gpu tensor, got {:?}", other),
898            }
899        });
900    }
901
902    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
903    #[test]
904    fn logical_bool_and_logical_inputs_passthrough() {
905        let res_bool = logical_builtin(Value::Bool(true), Vec::new()).expect("logical");
906        assert_eq!(res_bool, Value::Bool(true));
907
908        let logical = LogicalArray::new(vec![1, 0], vec![1, 2]).unwrap();
909        let res_array =
910            logical_builtin(Value::LogicalArray(logical.clone()), Vec::new()).expect("logical");
911        assert_eq!(res_array, Value::LogicalArray(logical));
912    }
913
914    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
915    #[test]
916    fn logical_empty_tensor_preserves_shape() {
917        let tensor = Tensor::new(Vec::new(), vec![0, 3]).unwrap();
918        let result = logical_builtin(Value::Tensor(tensor), Vec::new()).expect("logical");
919        match result {
920            Value::LogicalArray(array) => {
921                assert!(array.data.is_empty());
922                assert_eq!(array.shape, vec![0, 3]);
923            }
924            other => panic!("expected logical array, got {:?}", other),
925        }
926    }
927
928    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
929    #[test]
930    fn logical_integer_scalar() {
931        let res = logical_builtin(Value::Int(IntValue::I32(0)), Vec::new()).expect("logical");
932        assert_eq!(res, Value::Bool(false));
933
934        let res_nonzero =
935            logical_builtin(Value::Int(IntValue::I32(-5)), Vec::new()).expect("logical");
936        assert_eq!(res_nonzero, Value::Bool(true));
937    }
938
939    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
940    #[test]
941    #[cfg(feature = "wgpu")]
942    fn logical_wgpu_matches_cpu_conversion() {
943        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
944            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
945        ) else {
946            return;
947        };
948
949        let tensor = Tensor::new(vec![0.0, 2.0, -3.0, 1.0], vec![2, 2]).unwrap();
950        let cpu = logical_builtin(Value::Tensor(tensor.clone()), Vec::new()).unwrap();
951
952        let view = runmat_accelerate_api::HostTensorView {
953            data: &tensor.materialize_f64(),
954            shape: &tensor.shape,
955        };
956        let handle = provider.upload(&view).expect("upload");
957
958        let gpu_value = logical_builtin(Value::GpuTensor(handle), Vec::new()).unwrap();
959        let out_handle = match gpu_value {
960            Value::GpuTensor(ref h) => {
961                assert!(runmat_accelerate_api::handle_is_logical(h));
962                h.clone()
963            }
964            other => panic!("expected gpu tensor, got {other:?}"),
965        };
966
967        let gathered = test_support::gather(Value::GpuTensor(out_handle)).expect("gather");
968
969        let (expected, expected_shape): (Vec<f64>, Vec<usize>) = match cpu {
970            Value::LogicalArray(arr) => (
971                arr.data
972                    .iter()
973                    .map(|&b| if b != 0 { 1.0 } else { 0.0 })
974                    .collect(),
975                arr.shape.clone(),
976            ),
977            Value::Bool(flag) => (vec![if flag { 1.0 } else { 0.0 }], vec![1, 1]),
978            other => panic!("unexpected cpu result {other:?}"),
979        };
980
981        assert_eq!(gathered.shape, expected_shape);
982        assert_eq!(gathered.materialize_f64(), expected);
983    }
984}