1use 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}