1use 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
300pub 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 }
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
599pub 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}