Skip to main content

runmat_runtime/builtins/array/sorting_sets/
ismember.rs

1//! MATLAB-compatible `ismember` builtin with GPU-aware semantics for RunMat.
2
3use std::collections::HashMap;
4
5use runmat_accelerate_api::{
6    GpuTensorHandle, HostLogicalOwned, IsMemberOptions as ProviderIsMemberOptions, IsMemberResult,
7};
8use runmat_builtins::{
9    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
10    BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
11    BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind,
12    BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
13    BuiltinSignatureDescriptor,
14};
15use runmat_macros::runtime_builtin;
16use runmat_value::{
17    CharArray, ComplexStorage, ComplexTensor, IntValue, LogicalArray, NumericDType, NumericStorage,
18    StringArray, Tensor, Value,
19};
20
21use super::{float_order::SetFloat, type_resolvers::logical_output_type};
22use crate::build_runtime_error;
23use crate::builtins::common::gpu_helpers;
24use crate::builtins::common::spec::{
25    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
26    ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
27};
28use crate::builtins::common::tensor;
29use crate::builtins::math::elementwise::integer_cast::IntegerTarget;
30
31#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::array::sorting_sets::ismember")]
32pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
33    name: "ismember",
34    op_kind: GpuOpKind::Custom("ismember"),
35    supported_precisions: &[ScalarType::F32, ScalarType::F64],
36    broadcast: BroadcastSemantics::None,
37    provider_hooks: &[ProviderHook::Custom("ismember")],
38    constant_strategy: ConstantStrategy::InlineLiteral,
39    residency: ResidencyPolicy::NewHandle,
40    nan_mode: ReductionNaN::Include,
41    two_pass_threshold: None,
42    workgroup_size: None,
43    accepts_nan_mode: false,
44    notes: "Providers may supply dedicated membership kernels; exact typed fallback gathers when needed and restores logical membership plus double locations to the input owner.",
45};
46
47#[runmat_macros::register_fusion_spec(
48    builtin_path = "crate::builtins::array::sorting_sets::ismember"
49)]
50pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
51    name: "ismember",
52    shape: ShapeRequirements::Any,
53    constant_strategy: ConstantStrategy::InlineLiteral,
54    elementwise: None,
55    reduction: None,
56    emits_nan: false,
57    notes: "`ismember` materialises logical outputs and terminates fusion chains; upstream tensors are gathered when necessary.",
58};
59
60const BUILTIN_NAME: &str = "ismember";
61
62const ISMEMBER_OUTPUT_MASK: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
63    name: "tf",
64    ty: BuiltinParamType::LogicalArray,
65    arity: BuiltinParamArity::Required,
66    default: None,
67    description: "Membership mask over A.",
68}];
69
70const ISMEMBER_OUTPUT_MASK_LOC: [BuiltinParamDescriptor; 2] = [
71    BuiltinParamDescriptor {
72        name: "tf",
73        ty: BuiltinParamType::LogicalArray,
74        arity: BuiltinParamArity::Required,
75        default: None,
76        description: "Membership mask over A.",
77    },
78    BuiltinParamDescriptor {
79        name: "loc",
80        ty: BuiltinParamType::NumericArray,
81        arity: BuiltinParamArity::Required,
82        default: None,
83        description: "First-match indices into B for each element/row in A (0 when absent).",
84    },
85];
86
87const ISMEMBER_INPUTS_A_B: [BuiltinParamDescriptor; 2] = [
88    BuiltinParamDescriptor {
89        name: "A",
90        ty: BuiltinParamType::Any,
91        arity: BuiltinParamArity::Required,
92        default: None,
93        description: "Values or rows to query.",
94    },
95    BuiltinParamDescriptor {
96        name: "B",
97        ty: BuiltinParamType::Any,
98        arity: BuiltinParamArity::Required,
99        default: None,
100        description: "Reference set of values or rows.",
101    },
102];
103
104const ISMEMBER_INPUTS_A_B_OPTIONS: [BuiltinParamDescriptor; 3] = [
105    BuiltinParamDescriptor {
106        name: "A",
107        ty: BuiltinParamType::Any,
108        arity: BuiltinParamArity::Required,
109        default: None,
110        description: "Values or rows to query.",
111    },
112    BuiltinParamDescriptor {
113        name: "B",
114        ty: BuiltinParamType::Any,
115        arity: BuiltinParamArity::Required,
116        default: None,
117        description: "Reference set of values or rows.",
118    },
119    BuiltinParamDescriptor {
120        name: "option",
121        ty: BuiltinParamType::StringScalar,
122        arity: BuiltinParamArity::Variadic,
123        default: None,
124        description: "Option tokens: 'rows'.",
125    },
126];
127
128const ISMEMBER_SIGNATURES: [BuiltinSignatureDescriptor; 4] = [
129    BuiltinSignatureDescriptor {
130        label: "tf = ismember(A, B)",
131        inputs: &ISMEMBER_INPUTS_A_B,
132        outputs: &ISMEMBER_OUTPUT_MASK,
133    },
134    BuiltinSignatureDescriptor {
135        label: "tf = ismember(A, B, option...)",
136        inputs: &ISMEMBER_INPUTS_A_B_OPTIONS,
137        outputs: &ISMEMBER_OUTPUT_MASK,
138    },
139    BuiltinSignatureDescriptor {
140        label: "[tf, loc] = ismember(A, B)",
141        inputs: &ISMEMBER_INPUTS_A_B,
142        outputs: &ISMEMBER_OUTPUT_MASK_LOC,
143    },
144    BuiltinSignatureDescriptor {
145        label: "[tf, loc] = ismember(A, B, option...)",
146        inputs: &ISMEMBER_INPUTS_A_B_OPTIONS,
147        outputs: &ISMEMBER_OUTPUT_MASK_LOC,
148    },
149];
150
151const ISMEMBER_ERROR_LEGACY_OPTION_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
152    code: "RM.ISMEMBER.LEGACY_OPTION_UNSUPPORTED",
153    identifier: Some("RunMat:ismember:LegacyOptionUnsupported"),
154    when: "Legacy compatibility options are requested.",
155    message: "ismember: the 'legacy' behaviour is not supported",
156};
157
158const ISMEMBER_ERROR_UNKNOWN_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
159    code: "RM.ISMEMBER.UNKNOWN_OPTION",
160    identifier: Some("RunMat:ismember:UnknownOption"),
161    when: "An unsupported option token is provided.",
162    message: "ismember: unrecognised option",
163};
164
165const ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
166    code: "RM.ISMEMBER.ROWS_COLUMN_MISMATCH",
167    identifier: Some("RunMat:ismember:RowsColumnMismatch"),
168    when: "'rows' mode is used and column counts differ.",
169    message: "ismember: inputs must have the same number of columns when using 'rows'",
170};
171
172const ISMEMBER_ERROR_UNSUPPORTED_INPUT_TYPE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
173    code: "RM.ISMEMBER.UNSUPPORTED_INPUT_TYPE",
174    identifier: Some("RunMat:ismember:UnsupportedInputType"),
175    when: "Input classes or execution residency are unsupported.",
176    message: "ismember: unsupported input type",
177};
178
179const ISMEMBER_ERROR_NUMERIC_CLASS_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
180    code: "RM.ISMEMBER.NUMERIC_CLASS_MISMATCH",
181    identifier: Some("RunMat:ismember:NumericClassMismatch"),
182    when: "Numeric inputs have incompatible nondouble classes.",
183    message: "ismember: numeric inputs must have the same class, except double may be combined with one nondouble class",
184};
185
186const ISMEMBER_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
187    code: "RM.ISMEMBER.INVALID_ARGUMENT",
188    identifier: Some("RunMat:ismember:InvalidArgument"),
189    when: "Option arguments are not string-like where required.",
190    message: "ismember: expected string option arguments",
191};
192
193const ISMEMBER_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
194    code: "RM.ISMEMBER.INTERNAL",
195    identifier: Some("RunMat:ismember:Internal"),
196    when: "Internal conversion/allocation/provider decode fails.",
197    message: "ismember: internal operation failed",
198};
199
200const ISMEMBER_ERRORS: [BuiltinErrorDescriptor; 7] = [
201    ISMEMBER_ERROR_LEGACY_OPTION_UNSUPPORTED,
202    ISMEMBER_ERROR_UNKNOWN_OPTION,
203    ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH,
204    ISMEMBER_ERROR_UNSUPPORTED_INPUT_TYPE,
205    ISMEMBER_ERROR_NUMERIC_CLASS_MISMATCH,
206    ISMEMBER_ERROR_INVALID_ARGUMENT,
207    ISMEMBER_ERROR_INTERNAL,
208];
209
210const ISMEMBER_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
211    [BuiltinIntegerCapabilityDescriptor {
212        form: "[Lia, Locb] = ismember(integer_A, integer_B, options)",
213        inputs: &super::BINARY_SET_INTEGER_INPUTS,
214        computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
215        output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
216        overflow: BuiltinIntegerOverflowRule::NotApplicable,
217        backend: BuiltinIntegerBackendRule::GpuRestricted,
218        overload: BuiltinIntegerOverloadKind::Multiple,
219        notes: "Lia is logical and optional Locb is one-based double. Host supports all eight integer classes exactly; GPU supports integer classes through 32 bits, gathers typed fallback when needed, and restores both outputs to the owning provider.",
220    }];
221
222pub const ISMEMBER_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
223    signatures: &ISMEMBER_SIGNATURES,
224    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
225    completion_policy: BuiltinCompletionPolicy::Public,
226    errors: &ISMEMBER_ERRORS,
227};
228
229fn ismember_error_with(
230    error: &'static BuiltinErrorDescriptor,
231    message: impl Into<String>,
232) -> crate::RuntimeError {
233    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
234    if let Some(identifier) = error.identifier {
235        builder = builder.with_identifier(identifier);
236    }
237    builder.build()
238}
239
240fn ismember_error(error: &'static BuiltinErrorDescriptor) -> crate::RuntimeError {
241    ismember_error_with(error, error.message)
242}
243
244fn ismember_internal_error(message: impl Into<String>) -> crate::RuntimeError {
245    ismember_error_with(&ISMEMBER_ERROR_INTERNAL, message)
246}
247
248#[runtime_builtin(
249    name = "ismember",
250    category = "array/sorting_sets",
251    summary = "Identify array elements or rows that appear in another array while returning first-match indices.",
252    keywords = "ismember,membership,set,rows,indices,gpu",
253    accel = "array_construct",
254    sink = true,
255    type_resolver(logical_output_type),
256    descriptor(crate::builtins::array::sorting_sets::ismember::ISMEMBER_DESCRIPTOR),
257    integer_capabilities(ISMEMBER_INTEGER_CAPABILITIES),
258    builtin_path = "crate::builtins::array::sorting_sets::ismember"
259)]
260async fn ismember_builtin(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
261    if matches!(crate::output_count::current_output_count(), Some(n) if n > 2) {
262        return Err(ismember_error_with(
263            &ISMEMBER_ERROR_INVALID_ARGUMENT,
264            "ismember: too many output arguments; maximum is 2",
265        ));
266    }
267    let provider = super::set_output_provider(&a, &b);
268    let eval = evaluate(a, b, &rest).await?;
269    if let Some(out_count) = crate::output_count::current_output_count() {
270        if out_count == 0 {
271            return Ok(Value::OutputList(Vec::new()));
272        }
273        if out_count == 1 {
274            let outputs = super::restore_set_outputs(
275                provider,
276                BUILTIN_NAME,
277                vec![eval.into_mask_value()],
278                ismember_internal_error,
279            )?;
280            return Ok(Value::OutputList(outputs));
281        }
282        let (mask, loc) = eval.into_pair();
283        let outputs = super::restore_set_outputs(
284            provider,
285            BUILTIN_NAME,
286            vec![mask, loc],
287            ismember_internal_error,
288        )?;
289        return Ok(Value::OutputList(outputs));
290    }
291    let mut outputs = super::restore_set_outputs(
292        provider,
293        BUILTIN_NAME,
294        vec![eval.into_mask_value()],
295        ismember_internal_error,
296    )?;
297    Ok(outputs.pop().expect("ismember output"))
298}
299
300/// Evaluate the `ismember` builtin once and expose all outputs.
301pub async fn evaluate(
302    a: Value,
303    b: Value,
304    rest: &[Value],
305) -> crate::BuiltinResult<IsMemberEvaluation> {
306    crate::builtins::common::validation::reject_typed_complex_integer(&a, "ismember")?;
307    crate::builtins::common::validation::reject_typed_complex_integer(&b, "ismember")?;
308    let opts = parse_options(rest)?;
309    for value in [&a, &b] {
310        if let Value::GpuTensor(handle) = value {
311            if super::is_unsupported_set_gpu_integer(handle) {
312                return Err(ismember_error_with(
313                    &ISMEMBER_ERROR_UNSUPPORTED_INPUT_TYPE,
314                    "ismember: resident 64-bit integer inputs are not supported",
315                ));
316            }
317        }
318    }
319    match (a, b) {
320        (Value::GpuTensor(handle_a), Value::GpuTensor(handle_b)) => {
321            ismember_gpu_pair(handle_a, handle_b, &opts).await
322        }
323        (Value::GpuTensor(handle_a), other) => {
324            ismember_gpu_mixed(handle_a, other, &opts, true).await
325        }
326        (other, Value::GpuTensor(handle_b)) => {
327            ismember_gpu_mixed(handle_b, other, &opts, false).await
328        }
329        (left, right) => ismember_host(left, right, &opts),
330    }
331}
332
333#[derive(Debug, Clone, Copy)]
334struct IsMemberOptions {
335    rows: bool,
336}
337
338impl IsMemberOptions {
339    fn into_provider_options(self) -> ProviderIsMemberOptions {
340        ProviderIsMemberOptions { rows: self.rows }
341    }
342}
343
344fn parse_options(rest: &[Value]) -> crate::BuiltinResult<IsMemberOptions> {
345    let mut opts = IsMemberOptions { rows: false };
346    for arg in rest {
347        let text = tensor::value_to_string(arg)
348            .ok_or_else(|| ismember_error(&ISMEMBER_ERROR_INVALID_ARGUMENT))?;
349        let lowered = text.trim().to_ascii_lowercase();
350        match lowered.as_str() {
351            "rows" => opts.rows = true,
352            "legacy" | "r2012a" => {
353                return Err(ismember_error(&ISMEMBER_ERROR_LEGACY_OPTION_UNSUPPORTED))
354            }
355            other => {
356                return Err(ismember_error_with(
357                    &ISMEMBER_ERROR_UNKNOWN_OPTION,
358                    format!("ismember: unrecognised option '{other}'"),
359                ))
360            }
361        }
362    }
363    Ok(opts)
364}
365
366async fn ismember_gpu_pair(
367    handle_a: GpuTensorHandle,
368    handle_b: GpuTensorHandle,
369    opts: &IsMemberOptions,
370) -> crate::BuiltinResult<IsMemberEvaluation> {
371    if let Some(provider) = runmat_accelerate_api::provider_for_handle(&handle_a)
372        .or_else(runmat_accelerate_api::provider)
373    {
374        let provider_opts = opts.into_provider_options();
375        match provider
376            .ismember(&handle_a, &handle_b, &provider_opts)
377            .await
378        {
379            Ok(result) => return IsMemberEvaluation::from_provider_result(result),
380            Err(_) => {
381                // Fall back to host gather when the provider lacks an ismember implementation.
382            }
383        }
384    }
385    let tensor_a = gpu_helpers::gather_tensor_async(&handle_a).await?;
386    let tensor_b = gpu_helpers::gather_tensor_async(&handle_b).await?;
387    ismember_numeric_tensors(tensor_a, tensor_b, opts)
388}
389
390async fn ismember_gpu_mixed(
391    handle_gpu: GpuTensorHandle,
392    other: Value,
393    opts: &IsMemberOptions,
394    gpu_is_a: bool,
395) -> crate::BuiltinResult<IsMemberEvaluation> {
396    let tensor_gpu = gpu_helpers::gather_tensor_async(&handle_gpu).await?;
397    if gpu_is_a {
398        ismember_host(Value::Tensor(tensor_gpu), other, opts)
399    } else {
400        ismember_host(other, Value::Tensor(tensor_gpu), opts)
401    }
402}
403
404fn ismember_host(
405    a: Value,
406    b: Value,
407    opts: &IsMemberOptions,
408) -> crate::BuiltinResult<IsMemberEvaluation> {
409    match (a, b) {
410        (Value::ComplexTensor(at), Value::ComplexTensor(bt)) => ismember_complex(at, bt, opts.rows),
411        (Value::ComplexTensor(at), Value::Complex(re, im)) => {
412            let bt = ComplexTensor::new(vec![(re, im)], vec![1, 1])
413                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
414            ismember_complex(at, bt, opts.rows)
415        }
416        (Value::Complex(a_re, a_im), Value::ComplexTensor(bt)) => {
417            let at = ComplexTensor::new(vec![(a_re, a_im)], vec![1, 1])
418                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
419            ismember_complex(at, bt, opts.rows)
420        }
421        (Value::Complex(a_re, a_im), Value::Complex(b_re, b_im)) => {
422            let at = ComplexTensor::new(vec![(a_re, a_im)], vec![1, 1])
423                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
424            let bt = ComplexTensor::new(vec![(b_re, b_im)], vec![1, 1])
425                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
426            ismember_complex(at, bt, opts.rows)
427        }
428
429        (Value::CharArray(ac), Value::CharArray(bc)) => ismember_char(ac, bc, opts.rows),
430
431        (Value::StringArray(astring), Value::StringArray(bstring)) => {
432            ismember_string(astring, bstring, opts.rows)
433        }
434        (Value::StringArray(astring), Value::String(b)) => {
435            let bstring = StringArray::new(vec![b], vec![1, 1])
436                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
437            ismember_string(astring, bstring, opts.rows)
438        }
439        (Value::String(a), Value::StringArray(bstring)) => {
440            let astring = StringArray::new(vec![a], vec![1, 1])
441                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
442            ismember_string(astring, bstring, opts.rows)
443        }
444        (Value::String(a), Value::String(b)) => {
445            let astring = StringArray::new(vec![a], vec![1, 1])
446                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
447            let bstring = StringArray::new(vec![b], vec![1, 1])
448                .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
449            ismember_string(astring, bstring, opts.rows)
450        }
451
452        (left, right) => {
453            let tensor_a = tensor::value_into_tensor_for("ismember", left)
454                .map_err(|e| ismember_internal_error(e))?;
455            let tensor_b = tensor::value_into_tensor_for("ismember", right)
456                .map_err(|e| ismember_internal_error(e))?;
457            ismember_numeric_tensors(tensor_a, tensor_b, opts)
458        }
459    }
460}
461
462fn ismember_numeric_tensors(
463    a: Tensor,
464    b: Tensor,
465    opts: &IsMemberOptions,
466) -> crate::BuiltinResult<IsMemberEvaluation> {
467    let a_dtype = a.numeric_dtype();
468    let b_dtype = b.numeric_dtype();
469    if let (Some(a_storage), Some(b_storage)) = (a.integer_storage(), b.integer_storage()) {
470        if a_storage.class_name() == b_storage.class_name() {
471            return if opts.rows {
472                ismember_integer_rows(&a, &b)
473            } else {
474                ismember_integer_elements(&a, &b)
475            };
476        }
477        return Err(ismember_error(&ISMEMBER_ERROR_NUMERIC_CLASS_MISMATCH));
478    }
479    match (a.integer_storage(), b.integer_storage()) {
480        (Some(storage), None) if b_dtype == NumericDType::F64 => {
481            let target = IntegerTarget::from_storage(storage);
482            let b = target.cast_tensor(b).map_err(ismember_internal_error)?;
483            return ismember_numeric_tensors(a, b, opts);
484        }
485        (None, Some(storage)) if a_dtype == NumericDType::F64 => {
486            let target = IntegerTarget::from_storage(storage);
487            let a = target.cast_tensor(a).map_err(ismember_internal_error)?;
488            return ismember_numeric_tensors(a, b, opts);
489        }
490        _ => {}
491    }
492    if a_dtype != b_dtype && a_dtype != NumericDType::F64 && b_dtype != NumericDType::F64 {
493        return Err(ismember_error(&ISMEMBER_ERROR_NUMERIC_CLASS_MISMATCH));
494    }
495    let a_shape = a.shape.clone();
496    let b_shape = b.shape.clone();
497    let a_storage = a.into_numeric_storage().map_err(ismember_internal_error)?;
498    let b_storage = b.into_numeric_storage().map_err(ismember_internal_error)?;
499    match (a_storage, b_storage) {
500        (NumericStorage::F64(a), NumericStorage::F64(b)) => {
501            ismember_floating(a, a_shape, b, b_shape, opts.rows)
502        }
503        (NumericStorage::F32(a), NumericStorage::F32(b)) => {
504            ismember_floating(a, a_shape, b, b_shape, opts.rows)
505        }
506        (a, b) => ismember_promoted_f64(a, a_shape, b, b_shape, opts.rows),
507    }
508}
509
510fn ismember_promoted_f64(
511    a: NumericStorage,
512    a_shape: Vec<usize>,
513    b: NumericStorage,
514    b_shape: Vec<usize>,
515    rows: bool,
516) -> crate::BuiltinResult<IsMemberEvaluation> {
517    ismember_floating(
518        a.materialize_f64(),
519        a_shape,
520        b.materialize_f64(),
521        b_shape,
522        rows,
523    )
524}
525
526fn ismember_floating<T: SetFloat>(
527    a: Vec<T>,
528    a_shape: Vec<usize>,
529    b: Vec<T>,
530    b_shape: Vec<usize>,
531    rows: bool,
532) -> crate::BuiltinResult<IsMemberEvaluation> {
533    if rows {
534        ismember_floating_rows(a, a_shape, b, b_shape)
535    } else {
536        ismember_floating_elements(a, a_shape, b)
537    }
538}
539
540fn ismember_integer_elements(a: &Tensor, b: &Tensor) -> crate::BuiltinResult<IsMemberEvaluation> {
541    let a_values = a.integer_storage().expect("integer path").exact_values();
542    let b_values = b.integer_storage().expect("integer path").exact_values();
543    let mut map = HashMap::<IntValue, usize>::new();
544    for (index, value) in b_values.into_iter().enumerate() {
545        map.entry(value).or_insert(index + 1);
546    }
547    let mut mask = Vec::with_capacity(a_values.len());
548    let mut locations = Vec::with_capacity(a_values.len());
549    for value in a_values {
550        if let Some(&index) = map.get(&value) {
551            mask.push(1);
552            locations.push(index as f64);
553        } else {
554            mask.push(0);
555            locations.push(0.0);
556        }
557    }
558    let logical = LogicalArray::new(mask, a.shape.clone())
559        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
560    let locations = Tensor::new(locations, a.shape.clone())
561        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
562    Ok(IsMemberEvaluation::new(logical, locations))
563}
564
565fn ismember_integer_rows(a: &Tensor, b: &Tensor) -> crate::BuiltinResult<IsMemberEvaluation> {
566    let (rows_a, cols_a) = tensor_rows_cols(a, "ismember")?;
567    let (rows_b, cols_b) = tensor_rows_cols(b, "ismember")?;
568    if cols_a != cols_b {
569        return Err(ismember_error(&ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH));
570    }
571    let a_values = a.integer_storage().expect("integer path").exact_values();
572    let b_values = b.integer_storage().expect("integer path").exact_values();
573    let mut map = HashMap::<Vec<IntValue>, usize>::new();
574    for row in 0..rows_b {
575        let key: Vec<_> = (0..cols_b)
576            .map(|col| b_values[row + col * rows_b].clone())
577            .collect();
578        map.entry(key).or_insert(row + 1);
579    }
580    let mut mask = vec![0; rows_a];
581    let mut locations = vec![0.0; rows_a];
582    for row in 0..rows_a {
583        let key: Vec<_> = (0..cols_a)
584            .map(|col| a_values[row + col * rows_a].clone())
585            .collect();
586        if let Some(&index) = map.get(&key) {
587            mask[row] = 1;
588            locations[row] = index as f64;
589        }
590    }
591    let shape = vec![rows_a, 1];
592    let logical = LogicalArray::new(mask, shape.clone())
593        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
594    let locations = Tensor::new(locations, shape)
595        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
596    Ok(IsMemberEvaluation::new(logical, locations))
597}
598
599/// Helper exposed for acceleration providers handling numeric tensors on the host.
600pub fn ismember_numeric_from_tensors(
601    a: Tensor,
602    b: Tensor,
603    rows: bool,
604) -> crate::BuiltinResult<IsMemberEvaluation> {
605    let opts = IsMemberOptions { rows };
606    ismember_numeric_tensors(a, b, &opts)
607}
608
609#[cfg(test)]
610fn ismember_numeric_elements(a: Tensor, b: Tensor) -> crate::BuiltinResult<IsMemberEvaluation> {
611    ismember_numeric_tensors(a, b, &IsMemberOptions { rows: false })
612}
613
614fn ismember_floating_elements<T: SetFloat>(
615    a_values: Vec<T>,
616    a_shape: Vec<usize>,
617    b_values: Vec<T>,
618) -> crate::BuiltinResult<IsMemberEvaluation> {
619    let mut map: HashMap<u64, usize> = HashMap::new();
620    for (idx, &value) in b_values.iter().enumerate() {
621        map.entry(value.canonical_key()).or_insert(idx + 1);
622    }
623
624    let mut mask_data = Vec::<u8>::with_capacity(a_values.len());
625    let mut loc_data = Vec::<f64>::with_capacity(a_values.len());
626
627    for &value in a_values.iter() {
628        let key = value.canonical_key();
629        if let Some(&pos) = map.get(&key) {
630            mask_data.push(1);
631            loc_data.push(pos as f64);
632        } else {
633            mask_data.push(0);
634            loc_data.push(0.0);
635        }
636    }
637
638    let logical = LogicalArray::new(mask_data, a_shape.clone())
639        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
640    let loc_tensor = Tensor::new(loc_data, a_shape)
641        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
642    Ok(IsMemberEvaluation::new(logical, loc_tensor))
643}
644
645#[cfg(test)]
646fn ismember_numeric_rows(a: Tensor, b: Tensor) -> crate::BuiltinResult<IsMemberEvaluation> {
647    ismember_numeric_tensors(a, b, &IsMemberOptions { rows: true })
648}
649
650fn ismember_floating_rows<T: SetFloat>(
651    a_values: Vec<T>,
652    a_shape: Vec<usize>,
653    b_values: Vec<T>,
654    b_shape: Vec<usize>,
655) -> crate::BuiltinResult<IsMemberEvaluation> {
656    let (rows_a, cols_a) = shape_rows_cols(&a_shape, "ismember")?;
657    let (rows_b, cols_b) = shape_rows_cols(&b_shape, "ismember")?;
658    if cols_a != cols_b {
659        return Err(ismember_error(&ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH));
660    }
661    let mut map: HashMap<FloatingRowKey, usize> = HashMap::new();
662    for r in 0..rows_b {
663        let mut row_values = Vec::with_capacity(cols_b);
664        for c in 0..cols_b {
665            let idx = r + c * rows_b;
666            row_values.push(b_values[idx]);
667        }
668        let key = FloatingRowKey::from_slice(&row_values);
669        map.entry(key).or_insert(r + 1);
670    }
671
672    let mut mask_data = vec![0u8; rows_a];
673    let mut loc_data = vec![0.0f64; rows_a];
674
675    for r in 0..rows_a {
676        let mut row_values = Vec::with_capacity(cols_a);
677        for c in 0..cols_a {
678            let idx = r + c * rows_a;
679            row_values.push(a_values[idx]);
680        }
681        let key = FloatingRowKey::from_slice(&row_values);
682        if let Some(&pos) = map.get(&key) {
683            mask_data[r] = 1;
684            loc_data[r] = pos as f64;
685        }
686    }
687
688    let shape = vec![rows_a, 1];
689    let logical = LogicalArray::new(mask_data, shape.clone())
690        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
691    let loc_tensor = Tensor::new(loc_data, shape)
692        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
693    Ok(IsMemberEvaluation::new(logical, loc_tensor))
694}
695
696fn ismember_complex(
697    a: ComplexTensor,
698    b: ComplexTensor,
699    rows: bool,
700) -> crate::BuiltinResult<IsMemberEvaluation> {
701    let a_shape = a.shape.clone();
702    let b_shape = b.shape.clone();
703    match (a.into_complex_storage(), b.into_complex_storage()) {
704        (ComplexStorage::F64(a), ComplexStorage::F64(b)) => {
705            ismember_floating_complex(a, a_shape, b, b_shape, rows)
706        }
707        (ComplexStorage::F32(a), ComplexStorage::F32(b)) => {
708            ismember_floating_complex(a, a_shape, b, b_shape, rows)
709        }
710        (a, b) => ismember_promoted_complex_f64(a, a_shape, b, b_shape, rows),
711    }
712}
713
714fn ismember_promoted_complex_f64(
715    a: ComplexStorage,
716    a_shape: Vec<usize>,
717    b: ComplexStorage,
718    b_shape: Vec<usize>,
719    rows: bool,
720) -> crate::BuiltinResult<IsMemberEvaluation> {
721    ismember_floating_complex(
722        a.materialize_f64(),
723        a_shape,
724        b.materialize_f64(),
725        b_shape,
726        rows,
727    )
728}
729
730fn ismember_floating_complex<T: SetFloat>(
731    a: Vec<(T, T)>,
732    a_shape: Vec<usize>,
733    b: Vec<(T, T)>,
734    b_shape: Vec<usize>,
735    rows: bool,
736) -> crate::BuiltinResult<IsMemberEvaluation> {
737    if rows {
738        ismember_floating_complex_rows(a, a_shape, b, b_shape)
739    } else {
740        ismember_floating_complex_elements(a, a_shape, b)
741    }
742}
743
744#[cfg(test)]
745fn ismember_complex_elements(
746    a: ComplexTensor,
747    b: ComplexTensor,
748) -> crate::BuiltinResult<IsMemberEvaluation> {
749    ismember_complex(a, b, false)
750}
751
752fn ismember_floating_complex_elements<T: SetFloat>(
753    a: Vec<(T, T)>,
754    a_shape: Vec<usize>,
755    b: Vec<(T, T)>,
756) -> crate::BuiltinResult<IsMemberEvaluation> {
757    let mut map: HashMap<ComplexKey, usize> = HashMap::new();
758    for (idx, &value) in b.iter().enumerate() {
759        map.entry(ComplexKey::new(value)).or_insert(idx + 1);
760    }
761
762    let mut mask_data = Vec::<u8>::with_capacity(a.len());
763    let mut loc_data = Vec::<f64>::with_capacity(a.len());
764
765    for &value in &a {
766        let key = ComplexKey::new(value);
767        if let Some(&pos) = map.get(&key) {
768            mask_data.push(1);
769            loc_data.push(pos as f64);
770        } else {
771            mask_data.push(0);
772            loc_data.push(0.0);
773        }
774    }
775
776    let logical = LogicalArray::new(mask_data, a_shape.clone())
777        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
778    let loc_tensor = Tensor::new(loc_data, a_shape)
779        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
780    Ok(IsMemberEvaluation::new(logical, loc_tensor))
781}
782
783#[cfg(test)]
784fn ismember_complex_rows(
785    a: ComplexTensor,
786    b: ComplexTensor,
787) -> crate::BuiltinResult<IsMemberEvaluation> {
788    ismember_complex(a, b, true)
789}
790
791fn ismember_floating_complex_rows<T: SetFloat>(
792    a: Vec<(T, T)>,
793    a_shape: Vec<usize>,
794    b: Vec<(T, T)>,
795    b_shape: Vec<usize>,
796) -> crate::BuiltinResult<IsMemberEvaluation> {
797    let (rows_a, cols_a) = shape_rows_cols(&a_shape, "ismember")?;
798    let (rows_b, cols_b) = shape_rows_cols(&b_shape, "ismember")?;
799    if cols_a != cols_b {
800        return Err(ismember_error(&ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH).into());
801    }
802
803    let mut map: HashMap<Vec<ComplexKey>, usize> = HashMap::new();
804    for r in 0..rows_b {
805        let mut row_keys = Vec::with_capacity(cols_b);
806        for c in 0..cols_b {
807            let idx = r + c * rows_b;
808            row_keys.push(ComplexKey::new(b[idx]));
809        }
810        map.entry(row_keys).or_insert(r + 1);
811    }
812
813    let mut mask_data = vec![0u8; rows_a];
814    let mut loc_data = vec![0.0f64; rows_a];
815
816    for r in 0..rows_a {
817        let mut row_keys = Vec::with_capacity(cols_a);
818        for c in 0..cols_a {
819            let idx = r + c * rows_a;
820            row_keys.push(ComplexKey::new(a[idx]));
821        }
822        if let Some(&pos) = map.get(&row_keys) {
823            mask_data[r] = 1;
824            loc_data[r] = pos as f64;
825        }
826    }
827
828    let shape = vec![rows_a, 1];
829    let logical = LogicalArray::new(mask_data, shape.clone())
830        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
831    let loc_tensor = Tensor::new(loc_data, shape)
832        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
833    Ok(IsMemberEvaluation::new(logical, loc_tensor))
834}
835
836fn ismember_char(
837    a: CharArray,
838    b: CharArray,
839    rows: bool,
840) -> crate::BuiltinResult<IsMemberEvaluation> {
841    if rows {
842        ismember_char_rows(a, b)
843    } else {
844        ismember_char_elements(a, b)
845    }
846}
847
848fn ismember_char_elements(a: CharArray, b: CharArray) -> crate::BuiltinResult<IsMemberEvaluation> {
849    let rows_b = b.rows;
850    let cols_b = b.cols;
851    let mut map: HashMap<char, usize> = HashMap::new();
852
853    for col in 0..cols_b {
854        for row in 0..rows_b {
855            let data_idx = row * cols_b + col;
856            let ch = b.data[data_idx];
857            let linear_idx = row + col * rows_b;
858            map.entry(ch).or_insert(linear_idx + 1);
859        }
860    }
861
862    let rows_a = a.rows;
863    let cols_a = a.cols;
864    let mut mask_data = vec![0u8; rows_a * cols_a];
865    let mut loc_data = vec![0.0f64; rows_a * cols_a];
866
867    for col in 0..cols_a {
868        for row in 0..rows_a {
869            let data_idx = row * cols_a + col;
870            let ch = a.data[data_idx];
871            let linear_idx = row + col * rows_a;
872            if let Some(&pos) = map.get(&ch) {
873                mask_data[linear_idx] = 1;
874                loc_data[linear_idx] = pos as f64;
875            }
876        }
877    }
878
879    let shape = vec![rows_a, cols_a];
880    let logical = LogicalArray::new(mask_data, shape.clone())
881        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
882    let loc_tensor = Tensor::new(loc_data, shape)
883        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
884    Ok(IsMemberEvaluation::new(logical, loc_tensor))
885}
886
887fn ismember_char_rows(a: CharArray, b: CharArray) -> crate::BuiltinResult<IsMemberEvaluation> {
888    if a.cols != b.cols {
889        return Err(ismember_error(&ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH).into());
890    }
891
892    let rows_b = b.rows;
893    let cols = b.cols;
894    let mut map: HashMap<RowCharKey, usize> = HashMap::new();
895
896    for r in 0..rows_b {
897        let mut row_values = Vec::with_capacity(cols);
898        for c in 0..cols {
899            let idx = r * cols + c;
900            row_values.push(b.data[idx]);
901        }
902        let key = RowCharKey::from_slice(&row_values);
903        map.entry(key).or_insert(r + 1);
904    }
905
906    let rows_a = a.rows;
907    let mut mask_data = vec![0u8; rows_a];
908    let mut loc_data = vec![0.0f64; rows_a];
909
910    for r in 0..rows_a {
911        let mut row_values = Vec::with_capacity(cols);
912        for c in 0..cols {
913            let idx = r * cols + c;
914            row_values.push(a.data[idx]);
915        }
916        let key = RowCharKey::from_slice(&row_values);
917        if let Some(&pos) = map.get(&key) {
918            mask_data[r] = 1;
919            loc_data[r] = pos as f64;
920        }
921    }
922
923    let shape = vec![rows_a, 1];
924    let logical = LogicalArray::new(mask_data, shape.clone())
925        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
926    let loc_tensor = Tensor::new(loc_data, shape)
927        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
928    Ok(IsMemberEvaluation::new(logical, loc_tensor))
929}
930
931fn ismember_string(
932    a: StringArray,
933    b: StringArray,
934    rows: bool,
935) -> crate::BuiltinResult<IsMemberEvaluation> {
936    if rows {
937        ismember_string_rows(a, b)
938    } else {
939        ismember_string_elements(a, b)
940    }
941}
942
943fn ismember_string_elements(
944    a: StringArray,
945    b: StringArray,
946) -> crate::BuiltinResult<IsMemberEvaluation> {
947    let mut map: HashMap<String, usize> = HashMap::new();
948    for (idx, value) in b.data.iter().enumerate() {
949        map.entry(value.clone()).or_insert(idx + 1);
950    }
951
952    let mut mask_data = Vec::<u8>::with_capacity(a.data.len());
953    let mut loc_data = Vec::<f64>::with_capacity(a.data.len());
954
955    for value in &a.data {
956        if let Some(&pos) = map.get(value) {
957            mask_data.push(1);
958            loc_data.push(pos as f64);
959        } else {
960            mask_data.push(0);
961            loc_data.push(0.0);
962        }
963    }
964
965    let logical = LogicalArray::new(mask_data, a.shape.clone())
966        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
967    let loc_tensor = Tensor::new(loc_data, a.shape.clone())
968        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
969    Ok(IsMemberEvaluation::new(logical, loc_tensor))
970}
971
972fn ismember_string_rows(
973    a: StringArray,
974    b: StringArray,
975) -> crate::BuiltinResult<IsMemberEvaluation> {
976    if a.shape.len() != 2 || b.shape.len() != 2 {
977        return Err(ismember_internal_error(
978            "ismember: 'rows' option requires 2-D string arrays",
979        ));
980    }
981    if a.shape[1] != b.shape[1] {
982        return Err(ismember_error(&ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH).into());
983    }
984
985    let rows_a = a.shape[0];
986    let cols = a.shape[1];
987    let rows_b = b.shape[0];
988
989    let mut map: HashMap<RowStringKey, usize> = HashMap::new();
990    for r in 0..rows_b {
991        let mut row_values = Vec::with_capacity(cols);
992        for c in 0..cols {
993            let idx = r + c * rows_b;
994            row_values.push(b.data[idx].clone());
995        }
996        let key = RowStringKey(row_values);
997        map.entry(key).or_insert(r + 1);
998    }
999
1000    let mut mask_data = vec![0u8; rows_a];
1001    let mut loc_data = vec![0.0f64; rows_a];
1002
1003    for r in 0..rows_a {
1004        let mut row_values = Vec::with_capacity(cols);
1005        for c in 0..cols {
1006            let idx = r + c * rows_a;
1007            row_values.push(a.data[idx].clone());
1008        }
1009        let key = RowStringKey(row_values);
1010        if let Some(&pos) = map.get(&key) {
1011            mask_data[r] = 1;
1012            loc_data[r] = pos as f64;
1013        }
1014    }
1015
1016    let shape = vec![rows_a, 1];
1017    let logical = LogicalArray::new(mask_data, shape.clone())
1018        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
1019    let loc_tensor = Tensor::new(loc_data, shape)
1020        .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
1021    Ok(IsMemberEvaluation::new(logical, loc_tensor))
1022}
1023
1024fn tensor_rows_cols(t: &Tensor, name: &str) -> crate::BuiltinResult<(usize, usize)> {
1025    shape_rows_cols(&t.shape, name)
1026}
1027
1028fn shape_rows_cols(shape: &[usize], name: &str) -> crate::BuiltinResult<(usize, usize)> {
1029    match shape.len() {
1030        0 => Ok((1, 1)),
1031        1 => Ok((shape[0], 1)),
1032        2 => Ok((shape[0], shape[1])),
1033        _ => Err(ismember_internal_error(format!(
1034            "{name}: 'rows' option requires 2-D numeric matrices"
1035        ))
1036        .into()),
1037    }
1038}
1039
1040#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1041struct FloatingRowKey(Vec<u64>);
1042
1043impl FloatingRowKey {
1044    fn from_slice<T: SetFloat>(values: &[T]) -> Self {
1045        FloatingRowKey(values.iter().map(|&value| value.canonical_key()).collect())
1046    }
1047}
1048
1049#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
1050struct ComplexKey {
1051    re: u64,
1052    im: u64,
1053}
1054
1055impl ComplexKey {
1056    fn new<T: SetFloat>(value: (T, T)) -> Self {
1057        Self {
1058            re: value.0.canonical_key(),
1059            im: value.1.canonical_key(),
1060        }
1061    }
1062}
1063
1064#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1065struct RowCharKey(Vec<u32>);
1066
1067impl RowCharKey {
1068    fn from_slice(values: &[char]) -> Self {
1069        RowCharKey(values.iter().map(|&ch| ch as u32).collect())
1070    }
1071}
1072
1073#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1074struct RowStringKey(Vec<String>);
1075
1076#[derive(Debug, Clone)]
1077pub struct IsMemberEvaluation {
1078    mask: LogicalArray,
1079    loc: Tensor,
1080}
1081
1082impl IsMemberEvaluation {
1083    fn new(mask: LogicalArray, loc: Tensor) -> Self {
1084        Self { mask, loc }
1085    }
1086
1087    pub fn from_provider_result(result: IsMemberResult) -> crate::BuiltinResult<Self> {
1088        let mask = LogicalArray::new(result.mask.data, result.mask.shape)
1089            .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
1090        let loc = Tensor::new(result.loc.data, result.loc.shape)
1091            .map_err(|e| ismember_internal_error(format!("ismember: {e}")))?;
1092        Ok(IsMemberEvaluation::new(mask, loc))
1093    }
1094
1095    pub fn into_numeric_ismember_result(self) -> crate::BuiltinResult<IsMemberResult> {
1096        let IsMemberEvaluation { mask, loc } = self;
1097        Ok(IsMemberResult {
1098            mask: HostLogicalOwned {
1099                data: mask.data,
1100                shape: mask.shape,
1101            },
1102            loc: tensor::tensor_into_host_f64_owned(loc),
1103        })
1104    }
1105
1106    pub fn into_mask_value(self) -> Value {
1107        logical_array_into_value(self.mask)
1108    }
1109
1110    pub fn mask_value(&self) -> Value {
1111        logical_array_into_value(self.mask.clone())
1112    }
1113
1114    pub fn into_pair(self) -> (Value, Value) {
1115        let mask = logical_array_into_value(self.mask);
1116        let loc = tensor::tensor_into_value(self.loc);
1117        (mask, loc)
1118    }
1119
1120    pub fn loc_value(&self) -> Value {
1121        tensor::tensor_into_value(self.loc.clone())
1122    }
1123}
1124
1125fn logical_array_into_value(logical: LogicalArray) -> Value {
1126    if logical.data.len() == 1 {
1127        Value::Bool(logical.data[0] != 0)
1128    } else {
1129        Value::LogicalArray(logical)
1130    }
1131}
1132
1133#[cfg(test)]
1134pub(crate) mod tests {
1135    use super::*;
1136    use crate::builtins::common::test_support;
1137    use runmat_builtins::{ResolveContext, Type};
1138    use runmat_value::{IntegerStorage, Tensor};
1139
1140    #[cfg(feature = "wgpu")]
1141    use runmat_accelerate_api::HostTensorView;
1142
1143    fn evaluate_sync(
1144        a: Value,
1145        b: Value,
1146        rest: &[Value],
1147    ) -> crate::BuiltinResult<IsMemberEvaluation> {
1148        futures::executor::block_on(evaluate(a, b, rest))
1149    }
1150
1151    fn builtin_sync(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1152        futures::executor::block_on(ismember_builtin(a, b, rest))
1153    }
1154
1155    #[test]
1156    fn registered_builtin_restores_resident_outputs_and_rejects_excess_arity() {
1157        test_support::with_test_provider(|provider| {
1158            let left = Tensor::new_integer(IntegerStorage::I32(vec![7, 2, 9]), vec![3, 1]).unwrap();
1159            let right = Tensor::new_integer(IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
1160            let left =
1161                Value::GpuTensor(gpu_helpers::upload_tensor(provider, &left).expect("upload left"));
1162            let right = Value::GpuTensor(
1163                gpu_helpers::upload_tensor(provider, &right).expect("upload right"),
1164            );
1165
1166            {
1167                let _guard = crate::output_count::push_output_count(Some(2));
1168                let Value::OutputList(outputs) =
1169                    builtin_sync(left, right, Vec::new()).expect("resident ismember")
1170                else {
1171                    panic!("expected output list");
1172                };
1173                assert_eq!(outputs.len(), 2);
1174                let Value::GpuTensor(mask) = &outputs[0] else {
1175                    panic!("expected resident membership mask");
1176                };
1177                assert!(runmat_accelerate_api::handle_is_logical(mask));
1178                assert!(matches!(outputs[1], Value::GpuTensor(_)));
1179                assert_eq!(
1180                    test_support::gather(outputs[0].clone())
1181                        .expect("gather mask")
1182                        .materialize_f64(),
1183                    vec![1.0, 1.0, 0.0]
1184                );
1185                assert_eq!(
1186                    test_support::gather(outputs[1].clone())
1187                        .expect("gather locations")
1188                        .materialize_f64(),
1189                    vec![2.0, 1.0, 0.0]
1190                );
1191            }
1192
1193            let _guard = crate::output_count::push_output_count(Some(3));
1194            let err = builtin_sync(Value::Num(1.0), Value::Num(1.0), Vec::new())
1195                .expect_err("excess outputs must fail");
1196            assert_eq!(err.identifier(), ISMEMBER_ERROR_INVALID_ARGUMENT.identifier);
1197        });
1198    }
1199
1200    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1201    #[test]
1202    fn numeric_membership_basic() {
1203        let a = Tensor::new(vec![5.0, 7.0, 2.0, 7.0], vec![1, 4]).unwrap();
1204        let b = Tensor::new(vec![7.0, 9.0, 5.0], vec![1, 3]).unwrap();
1205        let eval = ismember_numeric_elements(a, b).expect("ismember");
1206        assert_eq!(eval.mask.data, vec![1, 1, 0, 1]);
1207        assert_eq!(eval.loc.materialize_f64(), vec![3.0, 1.0, 0.0, 1.0]);
1208    }
1209
1210    #[test]
1211    fn numeric_membership_uses_native_single_elements_and_rows() {
1212        let a = Tensor::from_f32(vec![1.0, 2.0, f32::NAN], vec![3, 1]).unwrap();
1213        let b = Tensor::from_f32(vec![2.0, f32::NAN], vec![2, 1]).unwrap();
1214        let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("single ismember");
1215        assert_eq!(eval.mask.data, vec![0, 1, 1]);
1216        assert_eq!(eval.loc.materialize_f64(), vec![0.0, 1.0, 2.0]);
1217
1218        let a = Tensor::from_f32(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
1219        let b = Tensor::from_f32(vec![3.0, 5.0, 4.0, 6.0], vec![2, 2]).unwrap();
1220        let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")])
1221            .expect("single row ismember");
1222        assert_eq!(eval.mask.data, vec![0, 1]);
1223        assert_eq!(eval.loc.materialize_f64(), vec![0.0, 1.0]);
1224    }
1225
1226    #[test]
1227    fn integer_membership_uses_exact_values_for_elements_and_rows() {
1228        let a = Tensor::new_integer(
1229            runmat_value::IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993]),
1230            vec![3, 1],
1231        )
1232        .expect("input");
1233        let b = Tensor::new_integer(
1234            runmat_value::IntegerStorage::U64(vec![0, u64::MAX]),
1235            vec![2, 1],
1236        )
1237        .expect("input");
1238        let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("ismember");
1239        assert_eq!(eval.mask.data, vec![1, 1, 0]);
1240        assert_eq!(eval.loc.materialize_f64(), vec![2.0, 1.0, 0.0]);
1241
1242        let a = Tensor::new_integer(
1243            runmat_value::IntegerStorage::I64(vec![i64::MAX, i64::MIN, 1, 2]),
1244            vec![2, 2],
1245        )
1246        .expect("input");
1247        let b = Tensor::new_integer(
1248            runmat_value::IntegerStorage::I64(vec![i64::MIN, 7, 2, 8]),
1249            vec![2, 2],
1250        )
1251        .expect("input");
1252        let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")])
1253            .expect("ismember rows");
1254        assert_eq!(eval.mask.data, vec![0, 1]);
1255        assert_eq!(eval.loc.materialize_f64(), vec![0.0, 1.0]);
1256    }
1257
1258    #[test]
1259    fn mixed_integer_membership_rejects_nondouble_class_mismatch() {
1260        let a = Tensor::new_integer(
1261            runmat_value::IntegerStorage::U16(vec![7, 2, 9, 7]),
1262            vec![4, 1],
1263        )
1264        .expect("input");
1265        let b = Tensor::new_integer(runmat_value::IntegerStorage::I32(vec![2, 7]), vec![2, 1])
1266            .expect("input");
1267
1268        let error = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
1269            .expect_err("mixed integer classes must reject");
1270        assert_eq!(
1271            error.identifier(),
1272            ISMEMBER_ERROR_NUMERIC_CLASS_MISMATCH.identifier
1273        );
1274    }
1275
1276    #[test]
1277    fn resident_integer_set_functions_use_exact_runtime_fallback_and_class_rules() {
1278        test_support::with_test_provider(|provider| {
1279            let left = Tensor::new_integer(
1280                runmat_value::IntegerStorage::I32(vec![7, 2, 9, 7]),
1281                vec![4, 1],
1282            )
1283            .unwrap();
1284            let right =
1285                Tensor::new_integer(runmat_value::IntegerStorage::I32(vec![2, 7]), vec![2, 1])
1286                    .unwrap();
1287            let left =
1288                Value::GpuTensor(gpu_helpers::upload_tensor(provider, &left).expect("upload left"));
1289            let right = Value::GpuTensor(
1290                gpu_helpers::upload_tensor(provider, &right).expect("upload right"),
1291            );
1292
1293            let member = futures::executor::block_on(evaluate(left.clone(), right.clone(), &[]))
1294                .expect("resident integer ismember");
1295            assert_eq!(member.mask.data, vec![1, 1, 0, 1]);
1296
1297            for (builtin, result) in [
1298                (
1299                    "intersect",
1300                    futures::executor::block_on(super::super::intersect::evaluate(
1301                        left.clone(),
1302                        right.clone(),
1303                        &[],
1304                    ))
1305                    .map(|eval| eval.values_value()),
1306                ),
1307                (
1308                    "union",
1309                    futures::executor::block_on(super::super::union::evaluate(
1310                        left.clone(),
1311                        right.clone(),
1312                        &[],
1313                    ))
1314                    .map(|eval| eval.values_value()),
1315                ),
1316                (
1317                    "setdiff",
1318                    futures::executor::block_on(super::super::setdiff::evaluate(
1319                        left.clone(),
1320                        right.clone(),
1321                        &[],
1322                    ))
1323                    .map(|eval| eval.values_value()),
1324                ),
1325                (
1326                    "setxor",
1327                    futures::executor::block_on(super::super::setxor::evaluate(
1328                        left.clone(),
1329                        right.clone(),
1330                        &[],
1331                    ))
1332                    .map(|eval| eval.values_value()),
1333                ),
1334            ] {
1335                let value = result.unwrap_or_else(|error| panic!("{builtin}: {error}"));
1336                let Value::Tensor(tensor) = value else {
1337                    panic!("{builtin}: expected integer tensor")
1338                };
1339                assert_eq!(
1340                    tensor.integer_storage().map(|storage| storage.class_name()),
1341                    Some("int32"),
1342                    "{builtin}"
1343                );
1344            }
1345
1346            let mismatched =
1347                Tensor::new_integer(runmat_value::IntegerStorage::I16(vec![2, 7]), vec![2, 1])
1348                    .unwrap();
1349            let mismatched =
1350                Value::GpuTensor(gpu_helpers::upload_tensor(provider, &mismatched).unwrap());
1351            for (builtin, error) in [
1352                (
1353                    "ismember",
1354                    futures::executor::block_on(evaluate(left.clone(), mismatched.clone(), &[]))
1355                        .expect_err("mismatch"),
1356                ),
1357                (
1358                    "intersect",
1359                    futures::executor::block_on(super::super::intersect::evaluate(
1360                        left.clone(),
1361                        mismatched.clone(),
1362                        &[],
1363                    ))
1364                    .expect_err("mismatch"),
1365                ),
1366                (
1367                    "union",
1368                    futures::executor::block_on(super::super::union::evaluate(
1369                        left.clone(),
1370                        mismatched.clone(),
1371                        &[],
1372                    ))
1373                    .expect_err("mismatch"),
1374                ),
1375                (
1376                    "setdiff",
1377                    futures::executor::block_on(super::super::setdiff::evaluate(
1378                        left.clone(),
1379                        mismatched.clone(),
1380                        &[],
1381                    ))
1382                    .expect_err("mismatch"),
1383                ),
1384                (
1385                    "setxor",
1386                    futures::executor::block_on(super::super::setxor::evaluate(
1387                        left.clone(),
1388                        mismatched,
1389                        &[],
1390                    ))
1391                    .expect_err("mismatch"),
1392                ),
1393            ] {
1394                assert!(
1395                    error
1396                        .identifier()
1397                        .is_some_and(|identifier| identifier.ends_with(":NumericClassMismatch")),
1398                    "{builtin}: {error}"
1399                );
1400            }
1401        });
1402    }
1403
1404    #[test]
1405    #[cfg(feature = "wgpu")]
1406    fn resident_i32_set_functions_use_exact_wgpu_runtime_fallback() {
1407        let _guard = test_support::accel_test_lock();
1408        let Ok(provider) = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1409            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1410        ) else {
1411            return;
1412        };
1413        let left = Tensor::new_integer(
1414            runmat_value::IntegerStorage::I32(vec![7, 2, 9, 7]),
1415            vec![4, 1],
1416        )
1417        .unwrap();
1418        let right =
1419            Tensor::new_integer(runmat_value::IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
1420        let left =
1421            Value::GpuTensor(gpu_helpers::upload_tensor(provider, &left).expect("upload left"));
1422        let right =
1423            Value::GpuTensor(gpu_helpers::upload_tensor(provider, &right).expect("upload right"));
1424
1425        let member = futures::executor::block_on(evaluate(left.clone(), right.clone(), &[]))
1426            .expect("wgpu integer ismember");
1427        assert_eq!(member.mask.data, vec![1, 1, 0, 1]);
1428
1429        for (builtin, result) in [
1430            (
1431                "intersect",
1432                futures::executor::block_on(super::super::intersect::evaluate(
1433                    left.clone(),
1434                    right.clone(),
1435                    &[],
1436                ))
1437                .map(|eval| eval.values_value()),
1438            ),
1439            (
1440                "union",
1441                futures::executor::block_on(super::super::union::evaluate(
1442                    left.clone(),
1443                    right.clone(),
1444                    &[],
1445                ))
1446                .map(|eval| eval.values_value()),
1447            ),
1448            (
1449                "setdiff",
1450                futures::executor::block_on(super::super::setdiff::evaluate(
1451                    left.clone(),
1452                    right.clone(),
1453                    &[],
1454                ))
1455                .map(|eval| eval.values_value()),
1456            ),
1457            (
1458                "setxor",
1459                futures::executor::block_on(super::super::setxor::evaluate(left, right, &[]))
1460                    .map(|eval| eval.values_value()),
1461            ),
1462        ] {
1463            let value = result.unwrap_or_else(|error| panic!("{builtin}: {error}"));
1464            let Value::Tensor(tensor) = value else {
1465                panic!("{builtin}: expected integer tensor")
1466            };
1467            assert_eq!(
1468                tensor.integer_storage().map(|storage| storage.class_name()),
1469                Some("int32"),
1470                "{builtin}"
1471            );
1472        }
1473    }
1474
1475    #[test]
1476    fn ismember_type_resolver_logical() {
1477        assert_eq!(
1478            logical_output_type(
1479                &[Type::tensor(), Type::tensor()],
1480                &ResolveContext::new(Vec::new()),
1481            ),
1482            Type::logical()
1483        );
1484    }
1485
1486    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1487    #[test]
1488    fn numeric_nan_membership() {
1489        let a = Tensor::new(vec![f64::NAN, 1.0], vec![1, 2]).unwrap();
1490        let b = Tensor::new(vec![f64::NAN, 2.0], vec![1, 2]).unwrap();
1491        let eval = ismember_numeric_elements(a, b).expect("ismember");
1492        assert_eq!(eval.mask.data, vec![1, 0]);
1493        assert_eq!(eval.loc.materialize_f64(), vec![1.0, 0.0]);
1494    }
1495
1496    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1497    #[test]
1498    fn numeric_rows_membership() {
1499        let a = Tensor::new(vec![1.0, 3.0, 1.0, 2.0, 4.0, 2.0], vec![3, 2]).unwrap();
1500        let b = Tensor::new(vec![3.0, 5.0, 1.0, 4.0, 6.0, 2.0], vec![3, 2]).unwrap();
1501        let eval = ismember_numeric_rows(a, b).expect("ismember");
1502        assert_eq!(eval.mask.data, vec![1, 1, 1]);
1503        assert_eq!(eval.loc.materialize_f64(), vec![3.0, 1.0, 3.0]);
1504        assert_eq!(eval.loc.shape, vec![3, 1]);
1505    }
1506
1507    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1508    #[test]
1509    fn complex_membership() {
1510        let a = ComplexTensor::new(vec![(1.0, 2.0), (0.0, 0.0)], vec![1, 2]).unwrap();
1511        let b = ComplexTensor::new(vec![(0.0, 0.0), (1.0, 2.0)], vec![1, 2]).unwrap();
1512        let eval = ismember_complex_elements(a, b).expect("ismember");
1513        assert_eq!(eval.mask.data, vec![1, 1]);
1514        assert_eq!(eval.loc.materialize_f64(), vec![2.0, 1.0]);
1515    }
1516
1517    #[test]
1518    fn complex_membership_uses_native_single_elements_and_rows() {
1519        let a = ComplexTensor::from_f32(vec![(1.0, 1.0), (2.0, 0.0)], vec![2, 1]).unwrap();
1520        let b = ComplexTensor::from_f32(vec![(2.0, 0.0), (1.0, 1.0)], vec![2, 1]).unwrap();
1521        let eval = evaluate_sync(Value::ComplexTensor(a), Value::ComplexTensor(b), &[])
1522            .expect("complex single ismember");
1523        assert_eq!(eval.mask.data, vec![1, 1]);
1524        assert_eq!(eval.loc.materialize_f64(), vec![2.0, 1.0]);
1525
1526        let a = ComplexTensor::from_f32(
1527            vec![(1.0, 0.0), (3.0, 0.0), (2.0, 1.0), (4.0, 1.0)],
1528            vec![2, 2],
1529        )
1530        .unwrap();
1531        let b = ComplexTensor::from_f32(
1532            vec![(3.0, 0.0), (5.0, 0.0), (4.0, 1.0), (6.0, 1.0)],
1533            vec![2, 2],
1534        )
1535        .unwrap();
1536        let eval = evaluate_sync(
1537            Value::ComplexTensor(a),
1538            Value::ComplexTensor(b),
1539            &[Value::from("rows")],
1540        )
1541        .expect("complex single row ismember");
1542        assert_eq!(eval.mask.data, vec![0, 1]);
1543        assert_eq!(eval.loc.materialize_f64(), vec![0.0, 1.0]);
1544    }
1545
1546    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1547    #[test]
1548    fn complex_rows_membership() {
1549        let a = ComplexTensor::new(
1550            vec![(1.0, 1.0), (3.0, 0.0), (2.0, 0.0), (4.0, 4.0)],
1551            vec![2, 2],
1552        )
1553        .unwrap();
1554        let b = ComplexTensor::new(
1555            vec![
1556                (1.0, 1.0),
1557                (5.0, 0.0),
1558                (3.0, 0.0),
1559                (2.0, 0.0),
1560                (6.0, 0.0),
1561                (4.0, 4.0),
1562            ],
1563            vec![3, 2],
1564        )
1565        .unwrap();
1566        let eval = ismember_complex_rows(a, b).expect("ismember");
1567        assert_eq!(eval.mask.data, vec![1, 1]);
1568        assert_eq!(eval.loc.materialize_f64(), vec![1.0, 3.0]);
1569    }
1570
1571    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1572    #[test]
1573    fn char_membership() {
1574        let a = CharArray::new(vec!['r', 'u', 'n', 'm'], 2, 2).unwrap();
1575        let b = CharArray::new(vec!['m', 'a', 'r', 'u'], 2, 2).unwrap();
1576        let eval = ismember_char_elements(a, b).expect("ismember");
1577        assert_eq!(eval.mask.data, vec![1, 0, 1, 1]);
1578        assert_eq!(eval.loc.materialize_f64(), vec![2.0, 0.0, 4.0, 1.0]);
1579    }
1580
1581    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1582    #[test]
1583    fn char_rows_membership() {
1584        let a = CharArray::new(vec!['m', 'a', 't', 'l'], 2, 2).unwrap();
1585        let b = CharArray::new(vec!['m', 'a', 'g', 'e', 't', 'l'], 3, 2).unwrap();
1586        let eval = ismember_char_rows(a, b).expect("ismember");
1587        assert_eq!(eval.mask.data, vec![1, 1]);
1588        assert_eq!(eval.loc.materialize_f64(), vec![1.0, 3.0]);
1589    }
1590
1591    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1592    #[test]
1593    fn string_membership() {
1594        let a = StringArray::new(
1595            vec![
1596                "apple".to_string(),
1597                "pear".to_string(),
1598                "banana".to_string(),
1599            ],
1600            vec![1, 3],
1601        )
1602        .unwrap();
1603        let b = StringArray::new(
1604            vec![
1605                "pear".to_string(),
1606                "orange".to_string(),
1607                "apple".to_string(),
1608            ],
1609            vec![1, 3],
1610        )
1611        .unwrap();
1612        let eval = ismember_string_elements(a, b).expect("ismember");
1613        assert_eq!(eval.mask.data, vec![1, 1, 0]);
1614        assert_eq!(eval.loc.materialize_f64(), vec![3.0, 1.0, 0.0]);
1615    }
1616
1617    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1618    #[test]
1619    fn string_rows_membership() {
1620        let a = StringArray::new(
1621            vec![
1622                "alpha".to_string(),
1623                "gamma".to_string(),
1624                "beta".to_string(),
1625                "delta".to_string(),
1626            ],
1627            vec![2, 2],
1628        )
1629        .unwrap();
1630        let b = StringArray::new(
1631            vec![
1632                "alpha".to_string(),
1633                "theta".to_string(),
1634                "gamma".to_string(),
1635                "beta".to_string(),
1636                "eta".to_string(),
1637                "delta".to_string(),
1638            ],
1639            vec![3, 2],
1640        )
1641        .unwrap();
1642        let eval = ismember_string_rows(a, b).expect("ismember");
1643        assert_eq!(eval.mask.data, vec![1, 1]);
1644        assert_eq!(eval.loc.materialize_f64(), vec![1.0, 3.0]);
1645    }
1646
1647    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1648    #[test]
1649    fn options_reject_legacy() {
1650        let err = parse_options(&[Value::from("legacy")]).unwrap_err();
1651        assert_eq!(
1652            err.identifier(),
1653            ISMEMBER_ERROR_LEGACY_OPTION_UNSUPPORTED.identifier
1654        );
1655    }
1656
1657    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1658    #[test]
1659    fn rejects_unknown_option() {
1660        let err =
1661            evaluate_sync(Value::Num(1.0), Value::Num(1.0), &[Value::from("stable")]).unwrap_err();
1662        assert_eq!(err.identifier(), ISMEMBER_ERROR_UNKNOWN_OPTION.identifier);
1663    }
1664
1665    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1666    #[test]
1667    fn ismember_runtime_numeric() {
1668        let a = Value::Tensor(Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap());
1669        let b = Value::Tensor(Tensor::new(vec![3.0, 1.0], vec![2, 1]).unwrap());
1670        let (mask, loc) = evaluate_sync(a, b, &[]).unwrap().into_pair();
1671        match mask {
1672            Value::LogicalArray(arr) => assert_eq!(arr.data, vec![1, 0, 1]),
1673            other => panic!("expected logical array, got {other:?}"),
1674        }
1675        match loc {
1676            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0, 0.0, 1.0]),
1677            other => panic!("expected tensor, got {other:?}"),
1678        }
1679    }
1680
1681    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1682    #[test]
1683    fn logical_inputs_promoted() {
1684        let a = Value::Bool(true);
1685        let logical_b =
1686            LogicalArray::new(vec![1, 0], vec![2, 1]).expect("logical array construction");
1687        let eval = evaluate_sync(a, Value::LogicalArray(logical_b), &[]).expect("ismember");
1688        assert_eq!(eval.mask_value(), Value::Bool(true));
1689        assert_eq!(eval.loc_value(), Value::Num(1.0));
1690    }
1691
1692    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1693    #[test]
1694    fn ismember_rows_shape_checks() {
1695        let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
1696        let b = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
1697        assert!(ismember_numeric_rows(a.clone(), b.clone()).is_ok());
1698        let bad = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
1699        let err = ismember_numeric_rows(a, bad).unwrap_err();
1700        assert_eq!(
1701            err.identifier(),
1702            ISMEMBER_ERROR_ROWS_COLUMN_MISMATCH.identifier
1703        );
1704    }
1705
1706    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1707    #[test]
1708    fn ismember_gpu_roundtrip() {
1709        test_support::with_test_provider(|provider| {
1710            let tensor = Tensor::new(vec![1.0, 4.0, 2.0, 4.0], vec![4, 1]).unwrap();
1711            let set = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
1712            let view_a = runmat_accelerate_api::HostTensorView {
1713                data: &tensor.materialize_f64(),
1714                shape: &tensor.shape,
1715            };
1716            let view_b = runmat_accelerate_api::HostTensorView {
1717                data: &set.materialize_f64(),
1718                shape: &set.shape,
1719            };
1720            let handle_a = provider.upload(&view_a).expect("upload a");
1721            let handle_b = provider.upload(&view_b).expect("upload b");
1722            let eval = evaluate_sync(Value::GpuTensor(handle_a), Value::GpuTensor(handle_b), &[])
1723                .expect("ismember");
1724            assert_eq!(eval.mask.data, vec![0, 1, 0, 1]);
1725            assert_eq!(eval.loc.materialize_f64(), vec![0.0, 1.0, 0.0, 1.0]);
1726        });
1727    }
1728
1729    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1730    #[test]
1731    fn ismember_gpu_rows_roundtrip() {
1732        test_support::with_test_provider(|provider| {
1733            let rows = Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
1734            let bank = Tensor::new(vec![1.0, 5.0, 3.0, 2.0, 6.0, 4.0], vec![3, 2]).unwrap();
1735            let view_a = runmat_accelerate_api::HostTensorView {
1736                data: &rows.materialize_f64(),
1737                shape: &rows.shape,
1738            };
1739            let view_b = runmat_accelerate_api::HostTensorView {
1740                data: &bank.materialize_f64(),
1741                shape: &bank.shape,
1742            };
1743            let handle_a = provider.upload(&view_a).expect("upload a");
1744            let handle_b = provider.upload(&view_b).expect("upload b");
1745            let eval = evaluate_sync(
1746                Value::GpuTensor(handle_a.clone()),
1747                Value::GpuTensor(handle_b.clone()),
1748                &[Value::from("rows")],
1749            )
1750            .expect("ismember");
1751            assert_eq!(eval.mask.data, vec![1, 1]);
1752            assert_eq!(eval.loc.materialize_f64(), vec![1.0, 3.0]);
1753            let _ = provider.free(&handle_a);
1754            let _ = provider.free(&handle_b);
1755        });
1756    }
1757
1758    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1759    #[test]
1760    #[cfg(feature = "wgpu")]
1761    fn ismember_wgpu_numeric_matches_cpu() {
1762        let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
1763            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
1764        );
1765
1766        let tensor = Tensor::new(vec![1.0, 4.0, 2.0, 4.0], vec![4, 1]).unwrap();
1767        let set = Tensor::new(vec![4.0, 5.0], vec![2, 1]).unwrap();
1768        let cpu_eval =
1769            ismember_numeric_from_tensors(tensor.clone(), set.clone(), false).expect("cpu");
1770
1771        let provider = runmat_accelerate_api::provider().expect("provider");
1772        let view_a = HostTensorView {
1773            data: &tensor.materialize_f64(),
1774            shape: &tensor.shape,
1775        };
1776        let view_b = HostTensorView {
1777            data: &set.materialize_f64(),
1778            shape: &set.shape,
1779        };
1780        let handle_a = provider.upload(&view_a).expect("upload a");
1781        let handle_b = provider.upload(&view_b).expect("upload b");
1782
1783        let eval = evaluate_sync(
1784            Value::GpuTensor(handle_a.clone()),
1785            Value::GpuTensor(handle_b.clone()),
1786            &[],
1787        )
1788        .expect("gpu evaluate");
1789        assert_eq!(eval.mask.data, cpu_eval.mask.data);
1790        assert_eq!(eval.loc.materialize_f64(), cpu_eval.loc.materialize_f64());
1791
1792        let _ = provider.free(&handle_a);
1793        let _ = provider.free(&handle_b);
1794
1795        let matrix = Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).unwrap();
1796        let bank = Tensor::new(vec![1.0, 7.0, 3.0, 2.0, 9.0, 4.0], vec![3, 2]).unwrap();
1797        let cpu_rows =
1798            ismember_numeric_from_tensors(matrix.clone(), bank.clone(), true).expect("cpu rows");
1799        let view_matrix = HostTensorView {
1800            data: &matrix.materialize_f64(),
1801            shape: &matrix.shape,
1802        };
1803        let view_bank = HostTensorView {
1804            data: &bank.materialize_f64(),
1805            shape: &bank.shape,
1806        };
1807        let handle_matrix = provider.upload(&view_matrix).expect("upload matrix");
1808        let handle_bank = provider.upload(&view_bank).expect("upload bank");
1809        let eval_rows = evaluate_sync(
1810            Value::GpuTensor(handle_matrix.clone()),
1811            Value::GpuTensor(handle_bank.clone()),
1812            &[Value::from("rows")],
1813        )
1814        .expect("gpu rows evaluate");
1815        assert_eq!(eval_rows.mask.data, cpu_rows.mask.data);
1816        assert_eq!(
1817            eval_rows.loc.materialize_f64(),
1818            cpu_rows.loc.materialize_f64()
1819        );
1820        let _ = provider.free(&handle_matrix);
1821        let _ = provider.free(&handle_bank);
1822    }
1823
1824    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1825    #[test]
1826    fn scalar_return_is_bool() {
1827        let a = Value::Tensor(Tensor::new(vec![7.0], vec![1, 1]).unwrap());
1828        let b = Value::Tensor(Tensor::new(vec![7.0], vec![1, 1]).unwrap());
1829        let mask = evaluate_sync(a, b, &[]).unwrap().into_mask_value();
1830        assert_eq!(mask, Value::Bool(true));
1831    }
1832
1833    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1834    #[test]
1835    fn parse_rows_option() {
1836        let opts = parse_options(&[Value::from("rows")]).unwrap();
1837        assert!(opts.rows);
1838    }
1839
1840    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1841    #[test]
1842    fn numeric_rows_with_nan() {
1843        let a = Tensor::new(vec![f64::NAN, 1.0], vec![2, 1]).unwrap();
1844        let b = Tensor::new(vec![f64::NAN, 2.0], vec![2, 1]).unwrap();
1845        let eval = ismember_numeric_rows(a, b).expect("ismember");
1846        assert_eq!(eval.mask.data, vec![1, 0]);
1847        assert_eq!(eval.loc.materialize_f64(), vec![1.0, 0.0]);
1848    }
1849}