1use std::cmp::Ordering;
9use std::collections::{HashMap, HashSet};
10
11use runmat_accelerate_api::GpuTensorHandle;
12use runmat_builtins::{
13 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
14 BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
15 BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind,
16 BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
17 BuiltinSignatureDescriptor,
18};
19use runmat_macros::runtime_builtin;
20use runmat_value::{
21 CharArray, ComplexStorage, ComplexTensor, IntValue, IntegerStorage, NumericDType,
22 NumericStorage, StringArray, Tensor, Value,
23};
24
25use super::{float_order::SetFloat, integer_order, type_resolvers::set_values_output_type};
26use crate::build_runtime_error;
27use crate::builtins::common::arg_tokens::tokens_from_values;
28use crate::builtins::common::gpu_helpers;
29use crate::builtins::common::random_args::complex_tensor_into_value;
30use crate::builtins::common::spec::{
31 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
32 ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
33};
34use crate::builtins::common::tensor;
35use crate::builtins::math::elementwise::integer_cast::IntegerTarget;
36
37#[runmat_macros::register_gpu_spec(
38 builtin_path = "crate::builtins::array::sorting_sets::intersect"
39)]
40pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
41 name: "intersect",
42 op_kind: GpuOpKind::Custom("intersect"),
43 supported_precisions: &[ScalarType::F32, ScalarType::F64],
44 broadcast: BroadcastSemantics::None,
45 provider_hooks: &[],
46 constant_strategy: ConstantStrategy::InlineLiteral,
47 residency: ResidencyPolicy::NewHandle,
48 nan_mode: ReductionNaN::Include,
49 two_pass_threshold: None,
50 workgroup_size: None,
51 accepts_nan_mode: true,
52 notes: "Exact typed fallback gathers when needed and restores intersection values plus double indices to the input owner.",
53};
54
55#[runmat_macros::register_fusion_spec(
56 builtin_path = "crate::builtins::array::sorting_sets::intersect"
57)]
58pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
59 name: "intersect",
60 shape: ShapeRequirements::Any,
61 constant_strategy: ConstantStrategy::InlineLiteral,
62 elementwise: None,
63 reduction: None,
64 emits_nan: true,
65 notes: "`intersect` materialises its inputs and terminates fusion chains; upstream GPU tensors are gathered when necessary.",
66};
67
68const BUILTIN_NAME: &str = "intersect";
69
70const INTERSECT_OUTPUT_C: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
71 name: "C",
72 ty: BuiltinParamType::Any,
73 arity: BuiltinParamArity::Required,
74 default: None,
75 description: "Intersection values or rows.",
76}];
77
78const INTERSECT_OUTPUT_C_IA: [BuiltinParamDescriptor; 2] = [
79 BuiltinParamDescriptor {
80 name: "C",
81 ty: BuiltinParamType::Any,
82 arity: BuiltinParamArity::Required,
83 default: None,
84 description: "Intersection values or rows.",
85 },
86 BuiltinParamDescriptor {
87 name: "ia",
88 ty: BuiltinParamType::NumericArray,
89 arity: BuiltinParamArity::Required,
90 default: None,
91 description: "Indices selecting matching elements/rows in A.",
92 },
93];
94
95const INTERSECT_OUTPUT_C_IA_IB: [BuiltinParamDescriptor; 3] = [
96 BuiltinParamDescriptor {
97 name: "C",
98 ty: BuiltinParamType::Any,
99 arity: BuiltinParamArity::Required,
100 default: None,
101 description: "Intersection values or rows.",
102 },
103 BuiltinParamDescriptor {
104 name: "ia",
105 ty: BuiltinParamType::NumericArray,
106 arity: BuiltinParamArity::Required,
107 default: None,
108 description: "Indices selecting matching elements/rows in A.",
109 },
110 BuiltinParamDescriptor {
111 name: "ib",
112 ty: BuiltinParamType::NumericArray,
113 arity: BuiltinParamArity::Required,
114 default: None,
115 description: "Indices selecting matching elements/rows in B.",
116 },
117];
118
119const INTERSECT_INPUTS_A_B: [BuiltinParamDescriptor; 2] = [
120 BuiltinParamDescriptor {
121 name: "A",
122 ty: BuiltinParamType::Any,
123 arity: BuiltinParamArity::Required,
124 default: None,
125 description: "First input array.",
126 },
127 BuiltinParamDescriptor {
128 name: "B",
129 ty: BuiltinParamType::Any,
130 arity: BuiltinParamArity::Required,
131 default: None,
132 description: "Second input array.",
133 },
134];
135
136const INTERSECT_INPUTS_A_B_OPTIONS: [BuiltinParamDescriptor; 3] = [
137 BuiltinParamDescriptor {
138 name: "A",
139 ty: BuiltinParamType::Any,
140 arity: BuiltinParamArity::Required,
141 default: None,
142 description: "First input array.",
143 },
144 BuiltinParamDescriptor {
145 name: "B",
146 ty: BuiltinParamType::Any,
147 arity: BuiltinParamArity::Required,
148 default: None,
149 description: "Second input array.",
150 },
151 BuiltinParamDescriptor {
152 name: "option",
153 ty: BuiltinParamType::StringScalar,
154 arity: BuiltinParamArity::Variadic,
155 default: None,
156 description: "Option tokens: 'rows'|'sorted'|'stable'.",
157 },
158];
159
160const INTERSECT_SIGNATURES: [BuiltinSignatureDescriptor; 6] = [
161 BuiltinSignatureDescriptor {
162 label: "C = intersect(A, B)",
163 inputs: &INTERSECT_INPUTS_A_B,
164 outputs: &INTERSECT_OUTPUT_C,
165 },
166 BuiltinSignatureDescriptor {
167 label: "C = intersect(A, B, option...)",
168 inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
169 outputs: &INTERSECT_OUTPUT_C,
170 },
171 BuiltinSignatureDescriptor {
172 label: "[C, ia] = intersect(A, B)",
173 inputs: &INTERSECT_INPUTS_A_B,
174 outputs: &INTERSECT_OUTPUT_C_IA,
175 },
176 BuiltinSignatureDescriptor {
177 label: "[C, ia] = intersect(A, B, option...)",
178 inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
179 outputs: &INTERSECT_OUTPUT_C_IA,
180 },
181 BuiltinSignatureDescriptor {
182 label: "[C, ia, ib] = intersect(A, B)",
183 inputs: &INTERSECT_INPUTS_A_B,
184 outputs: &INTERSECT_OUTPUT_C_IA_IB,
185 },
186 BuiltinSignatureDescriptor {
187 label: "[C, ia, ib] = intersect(A, B, option...)",
188 inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
189 outputs: &INTERSECT_OUTPUT_C_IA_IB,
190 },
191];
192
193const INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
194 code: "RM.INTERSECT.LEGACY_OPTION_UNSUPPORTED",
195 identifier: Some("RunMat:intersect:LegacyOptionUnsupported"),
196 when: "Legacy compatibility options are requested.",
197 message: "intersect: the 'legacy' behaviour is not supported",
198};
199
200const INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
201 code: "RM.INTERSECT.CONFLICTING_ORDER_OPTIONS",
202 identifier: Some("RunMat:intersect:ConflictingOrderOptions"),
203 when: "Both 'sorted' and 'stable' options are provided.",
204 message: "intersect: cannot combine 'sorted' with 'stable'",
205};
206
207const INTERSECT_ERROR_UNKNOWN_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
208 code: "RM.INTERSECT.UNKNOWN_OPTION",
209 identifier: Some("RunMat:intersect:UnknownOption"),
210 when: "An unsupported option token is provided.",
211 message: "intersect: unrecognised option",
212};
213
214const INTERSECT_ERROR_ROWS_COLUMN_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
215 code: "RM.INTERSECT.ROWS_COLUMN_MISMATCH",
216 identifier: Some("RunMat:intersect:RowsColumnMismatch"),
217 when: "'rows' mode is used and column counts differ.",
218 message: "intersect: inputs must have the same number of columns when using 'rows'",
219};
220
221const INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
222 code: "RM.INTERSECT.UNSUPPORTED_INPUT_TYPE",
223 identifier: Some("RunMat:intersect:UnsupportedInputType"),
224 when: "Input values cannot be converted into supported intersect domains.",
225 message: "intersect: unsupported input type",
226};
227
228const INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
229 code: "RM.INTERSECT.NUMERIC_CLASS_MISMATCH",
230 identifier: Some("RunMat:intersect:NumericClassMismatch"),
231 when: "Numeric inputs have incompatible nondouble classes.",
232 message: "intersect: numeric inputs must have the same class, except double may be combined with one nondouble class",
233};
234
235const INTERSECT_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
236 code: "RM.INTERSECT.INVALID_ARGUMENT",
237 identifier: Some("RunMat:intersect:InvalidArgument"),
238 when: "Option arguments are not string-like where required.",
239 message: "intersect: expected string option arguments",
240};
241
242const INTERSECT_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
243 code: "RM.INTERSECT.INTERNAL",
244 identifier: Some("RunMat:intersect:Internal"),
245 when: "Internal conversion/allocation/provider decode fails.",
246 message: "intersect: internal operation failed",
247};
248
249const INTERSECT_ERRORS: [BuiltinErrorDescriptor; 8] = [
250 INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED,
251 INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS,
252 INTERSECT_ERROR_UNKNOWN_OPTION,
253 INTERSECT_ERROR_ROWS_COLUMN_MISMATCH,
254 INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE,
255 INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH,
256 INTERSECT_ERROR_INVALID_ARGUMENT,
257 INTERSECT_ERROR_INTERNAL,
258];
259
260const INTERSECT_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
261 [BuiltinIntegerCapabilityDescriptor {
262 form: "[C, ia, ib] = intersect(integer_A, integer_B, options)",
263 inputs: &super::BINARY_SET_INTEGER_INPUTS,
264 computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
265 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
266 overflow: BuiltinIntegerOverflowRule::NotApplicable,
267 backend: BuiltinIntegerBackendRule::GpuRestricted,
268 overload: BuiltinIntegerOverloadKind::Multiple,
269 notes: "C preserves the common nondouble integer class, including when paired with double; ia and ib are one-based double. GPU supports integer classes through 32 bits and restores outputs after typed fallback.",
270 }];
271
272pub const INTERSECT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
273 signatures: &INTERSECT_SIGNATURES,
274 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
275 completion_policy: BuiltinCompletionPolicy::Public,
276 errors: &INTERSECT_ERRORS,
277};
278
279fn intersect_error_with(
280 error: &'static BuiltinErrorDescriptor,
281 message: impl Into<String>,
282) -> crate::RuntimeError {
283 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
284 if let Some(identifier) = error.identifier {
285 builder = builder.with_identifier(identifier);
286 }
287 builder.build()
288}
289
290fn intersect_error(error: &'static BuiltinErrorDescriptor) -> crate::RuntimeError {
291 intersect_error_with(error, error.message)
292}
293
294fn intersect_internal_error(message: impl Into<String>) -> crate::RuntimeError {
295 intersect_error_with(&INTERSECT_ERROR_INTERNAL, message)
296}
297
298#[runtime_builtin(
299 name = "intersect",
300 category = "array/sorting_sets",
301 summary = "Return common elements or rows across arrays with index outputs.",
302 keywords = "intersect,set,stable,rows,indices,gpu",
303 accel = "array_construct",
304 sink = true,
305 type_resolver(set_values_output_type),
306 descriptor(crate::builtins::array::sorting_sets::intersect::INTERSECT_DESCRIPTOR),
307 integer_capabilities(INTERSECT_INTEGER_CAPABILITIES),
308 builtin_path = "crate::builtins::array::sorting_sets::intersect"
309)]
310async fn intersect_builtin(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
311 if matches!(crate::output_count::current_output_count(), Some(n) if n > 3) {
312 return Err(intersect_error_with(
313 &INTERSECT_ERROR_INVALID_ARGUMENT,
314 "intersect: too many output arguments; maximum is 3",
315 ));
316 }
317 let provider = super::set_output_provider(&a, &b);
318 let eval = evaluate(a, b, &rest).await?;
319 if let Some(out_count) = crate::output_count::current_output_count() {
320 if out_count == 0 {
321 return Ok(Value::OutputList(Vec::new()));
322 }
323 if out_count == 1 {
324 let outputs = super::restore_set_outputs(
325 provider,
326 BUILTIN_NAME,
327 vec![eval.into_values_value()],
328 intersect_internal_error,
329 )?;
330 return Ok(Value::OutputList(outputs));
331 }
332 if out_count == 2 {
333 let (values, ia) = eval.into_pair();
334 let outputs = super::restore_set_outputs(
335 provider,
336 BUILTIN_NAME,
337 vec![values, ia],
338 intersect_internal_error,
339 )?;
340 return Ok(Value::OutputList(outputs));
341 }
342 let (values, ia, ib) = eval.into_triple();
343 let outputs = super::restore_set_outputs(
344 provider,
345 BUILTIN_NAME,
346 vec![values, ia, ib],
347 intersect_internal_error,
348 )?;
349 return Ok(Value::OutputList(outputs));
350 }
351 let mut outputs = super::restore_set_outputs(
352 provider,
353 BUILTIN_NAME,
354 vec![eval.into_values_value()],
355 intersect_internal_error,
356 )?;
357 Ok(outputs.pop().expect("intersect output"))
358}
359
360pub async fn evaluate(
362 a: Value,
363 b: Value,
364 rest: &[Value],
365) -> crate::BuiltinResult<IntersectEvaluation> {
366 crate::builtins::common::validation::reject_typed_complex_integer(&a, "intersect")?;
367 crate::builtins::common::validation::reject_typed_complex_integer(&b, "intersect")?;
368 let opts = parse_options(rest)?;
369 for value in [&a, &b] {
370 if let Value::GpuTensor(handle) = value {
371 if super::is_unsupported_set_gpu_integer(handle) {
372 return Err(intersect_error_with(
373 &INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE,
374 "intersect: resident 64-bit integer inputs are not supported",
375 ));
376 }
377 }
378 }
379 match (a, b) {
380 (Value::GpuTensor(handle_a), Value::GpuTensor(handle_b)) => {
381 intersect_gpu_pair(handle_a, handle_b, &opts).await
382 }
383 (Value::GpuTensor(handle_a), other) => {
384 intersect_gpu_mixed(handle_a, other, &opts, true).await
385 }
386 (other, Value::GpuTensor(handle_b)) => {
387 intersect_gpu_mixed(handle_b, other, &opts, false).await
388 }
389 (left, right) => intersect_host(left, right, &opts),
390 }
391}
392
393#[derive(Debug, Clone, Copy, PartialEq, Eq)]
394enum IntersectOrder {
395 Sorted,
396 Stable,
397}
398
399#[derive(Debug, Clone)]
400struct IntersectOptions {
401 rows: bool,
402 order: IntersectOrder,
403}
404
405fn parse_options(rest: &[Value]) -> crate::BuiltinResult<IntersectOptions> {
406 let mut opts = IntersectOptions {
407 rows: false,
408 order: IntersectOrder::Sorted,
409 };
410 let mut seen_order: Option<IntersectOrder> = None;
411
412 let tokens = tokens_from_values(rest);
413 for (arg, token) in rest.iter().zip(tokens.iter()) {
414 let text = match token {
415 crate::builtins::common::arg_tokens::ArgToken::String(text) => text.as_str(),
416 _ => {
417 let text = tensor::value_to_string(arg)
418 .ok_or_else(|| intersect_error(&INTERSECT_ERROR_INVALID_ARGUMENT))?;
419 let lowered = text.trim().to_ascii_lowercase();
420 parse_intersect_option(&mut opts, &mut seen_order, &lowered)?;
421 continue;
422 }
423 };
424 parse_intersect_option(&mut opts, &mut seen_order, text)?;
425 }
426
427 Ok(opts)
428}
429
430fn parse_intersect_option(
431 opts: &mut IntersectOptions,
432 seen_order: &mut Option<IntersectOrder>,
433 lowered: &str,
434) -> crate::BuiltinResult<()> {
435 match lowered {
436 "rows" => opts.rows = true,
437 "sorted" => {
438 if let Some(prev) = seen_order {
439 if *prev != IntersectOrder::Sorted {
440 return Err(intersect_error(&INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS));
441 }
442 }
443 *seen_order = Some(IntersectOrder::Sorted);
444 opts.order = IntersectOrder::Sorted;
445 }
446 "stable" => {
447 if let Some(prev) = seen_order {
448 if *prev != IntersectOrder::Stable {
449 return Err(intersect_error(&INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS));
450 }
451 }
452 *seen_order = Some(IntersectOrder::Stable);
453 opts.order = IntersectOrder::Stable;
454 }
455 "legacy" | "r2012a" => {
456 return Err(intersect_error(&INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED));
457 }
458 other => {
459 return Err(intersect_error_with(
460 &INTERSECT_ERROR_UNKNOWN_OPTION,
461 format!("intersect: unrecognised option '{other}'"),
462 ))
463 }
464 }
465 Ok(())
466}
467
468async fn intersect_gpu_pair(
469 handle_a: GpuTensorHandle,
470 handle_b: GpuTensorHandle,
471 opts: &IntersectOptions,
472) -> crate::BuiltinResult<IntersectEvaluation> {
473 let tensor_a = gpu_helpers::gather_tensor_async(&handle_a).await?;
474 let tensor_b = gpu_helpers::gather_tensor_async(&handle_b).await?;
475 intersect_numeric(tensor_a, tensor_b, opts)
476}
477
478async fn intersect_gpu_mixed(
479 handle_gpu: GpuTensorHandle,
480 other: Value,
481 opts: &IntersectOptions,
482 gpu_is_a: bool,
483) -> crate::BuiltinResult<IntersectEvaluation> {
484 let tensor_gpu = gpu_helpers::gather_tensor_async(&handle_gpu).await?;
485 let tensor_other = tensor::value_into_tensor_for("intersect", other)
486 .map_err(|e| intersect_internal_error(e))?;
487 if gpu_is_a {
488 intersect_numeric(tensor_gpu, tensor_other, opts)
489 } else {
490 intersect_numeric(tensor_other, tensor_gpu, opts)
491 }
492}
493
494fn intersect_host(
495 a: Value,
496 b: Value,
497 opts: &IntersectOptions,
498) -> crate::BuiltinResult<IntersectEvaluation> {
499 match (a, b) {
500 (Value::ComplexTensor(at), Value::ComplexTensor(bt)) => intersect_complex(at, bt, opts),
501 (Value::ComplexTensor(at), Value::Complex(re, im)) => {
502 let bt = scalar_complex_tensor(re, im)?;
503 intersect_complex(at, bt, opts)
504 }
505 (Value::Complex(re, im), Value::ComplexTensor(bt)) => {
506 let at = scalar_complex_tensor(re, im)?;
507 intersect_complex(at, bt, opts)
508 }
509 (Value::Complex(a_re, a_im), Value::Complex(b_re, b_im)) => {
510 let at = scalar_complex_tensor(a_re, a_im)?;
511 let bt = scalar_complex_tensor(b_re, b_im)?;
512 intersect_complex(at, bt, opts)
513 }
514 (Value::ComplexTensor(at), other) => {
515 let bt = value_into_complex_tensor(other)?;
516 intersect_complex(at, bt, opts)
517 }
518 (other, Value::ComplexTensor(bt)) => {
519 let at = value_into_complex_tensor(other)?;
520 intersect_complex(at, bt, opts)
521 }
522 (Value::Complex(re, im), other) => {
523 let at = scalar_complex_tensor(re, im)?;
524 let bt = value_into_complex_tensor(other)?;
525 intersect_complex(at, bt, opts)
526 }
527 (other, Value::Complex(re, im)) => {
528 let at = value_into_complex_tensor(other)?;
529 let bt = scalar_complex_tensor(re, im)?;
530 intersect_complex(at, bt, opts)
531 }
532
533 (Value::CharArray(ac), Value::CharArray(bc)) => intersect_char(ac, bc, opts),
534
535 (Value::StringArray(astring), Value::StringArray(bstring)) => {
536 intersect_string(astring, bstring, opts)
537 }
538 (Value::StringArray(astring), Value::String(b)) => {
539 let bstring = StringArray::new(vec![b], vec![1, 1])
540 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
541 intersect_string(astring, bstring, opts)
542 }
543 (Value::String(a), Value::StringArray(bstring)) => {
544 let astring = StringArray::new(vec![a], vec![1, 1])
545 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
546 intersect_string(astring, bstring, opts)
547 }
548 (Value::String(a), Value::String(b)) => {
549 let astring = StringArray::new(vec![a], vec![1, 1])
550 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
551 let bstring = StringArray::new(vec![b], vec![1, 1])
552 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
553 intersect_string(astring, bstring, opts)
554 }
555
556 (left, right) => {
557 let tensor_a = tensor::value_into_tensor_for("intersect", left)
558 .map_err(|e| intersect_error_with(&INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
559 let tensor_b = tensor::value_into_tensor_for("intersect", right)
560 .map_err(|e| intersect_error_with(&INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
561 intersect_numeric(tensor_a, tensor_b, opts)
562 }
563 }
564}
565
566fn intersect_numeric(
567 a: Tensor,
568 b: Tensor,
569 opts: &IntersectOptions,
570) -> crate::BuiltinResult<IntersectEvaluation> {
571 let a_dtype = a.numeric_dtype();
572 let b_dtype = b.numeric_dtype();
573 if let (Some(a_storage), Some(b_storage)) = (a.integer_storage(), b.integer_storage()) {
574 if a_storage.class_name() == b_storage.class_name() {
575 return if opts.rows {
576 intersect_integer_rows(a_storage, a.shape.clone(), b_storage, b.shape.clone(), opts)
577 } else {
578 intersect_integer_elements(a_storage, b_storage, opts)
579 };
580 }
581 return Err(intersect_error(&INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH));
582 }
583 match (a.integer_storage(), b.integer_storage()) {
584 (Some(storage), None) if b_dtype == NumericDType::F64 => {
585 let target = IntegerTarget::from_storage(storage);
586 let b = target.cast_tensor(b).map_err(intersect_internal_error)?;
587 return intersect_numeric(a, b, opts);
588 }
589 (None, Some(storage)) if a_dtype == NumericDType::F64 => {
590 let target = IntegerTarget::from_storage(storage);
591 let a = target.cast_tensor(a).map_err(intersect_internal_error)?;
592 return intersect_numeric(a, b, opts);
593 }
594 _ => {}
595 }
596 if a_dtype != b_dtype && a_dtype != NumericDType::F64 && b_dtype != NumericDType::F64 {
597 return Err(intersect_error(&INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH));
598 }
599 let a_shape = a.shape.clone();
600 let b_shape = b.shape.clone();
601 let a_storage = a
602 .into_numeric_storage()
603 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
604 let b_storage = b
605 .into_numeric_storage()
606 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
607 match (a_storage, b_storage) {
608 (NumericStorage::F64(a), NumericStorage::F64(b)) => {
609 intersect_floating(a, a_shape, b, b_shape, opts)
610 }
611 (NumericStorage::F32(a), NumericStorage::F32(b)) => {
612 intersect_floating(a, a_shape, b, b_shape, opts)
613 }
614 (a, b) => intersect_promoted_f64(a, a_shape, b, b_shape, opts),
615 }
616}
617
618fn intersect_promoted_f64(
619 a: NumericStorage,
620 a_shape: Vec<usize>,
621 b: NumericStorage,
622 b_shape: Vec<usize>,
623 opts: &IntersectOptions,
624) -> crate::BuiltinResult<IntersectEvaluation> {
625 intersect_floating(
626 a.materialize_f64(),
627 a_shape,
628 b.materialize_f64(),
629 b_shape,
630 opts,
631 )
632}
633
634fn intersect_floating<T: SetFloat>(
635 a: Vec<T>,
636 a_shape: Vec<usize>,
637 b: Vec<T>,
638 b_shape: Vec<usize>,
639 opts: &IntersectOptions,
640) -> crate::BuiltinResult<IntersectEvaluation> {
641 if opts.rows {
642 intersect_floating_rows(a, a_shape, b, b_shape, opts)
643 } else {
644 intersect_floating_elements(a, b, opts)
645 }
646}
647
648fn intersect_integer_elements(
649 a: &IntegerStorage,
650 b: &IntegerStorage,
651 opts: &IntersectOptions,
652) -> crate::BuiltinResult<IntersectEvaluation> {
653 let mut b_map = HashMap::<IntValue, usize>::new();
654 for (index, value) in b.exact_values().into_iter().enumerate() {
655 b_map.entry(value).or_insert(index);
656 }
657 let mut seen = HashSet::<IntValue>::new();
658 let mut entries = Vec::<IntegerIntersectEntry>::new();
659 for (a_index, value) in a.exact_values().into_iter().enumerate() {
660 if seen.contains(&value) {
661 continue;
662 }
663 if let Some(&b_index) = b_map.get(&value) {
664 let order_rank = entries.len();
665 entries.push(IntegerIntersectEntry {
666 value: value.clone(),
667 a_index,
668 b_index,
669 order_rank,
670 });
671 seen.insert(value);
672 }
673 }
674 assemble_integer_intersect(entries, a, opts)
675}
676
677fn intersect_integer_rows(
678 a_storage: &IntegerStorage,
679 a_shape: Vec<usize>,
680 b_storage: &IntegerStorage,
681 b_shape: Vec<usize>,
682 opts: &IntersectOptions,
683) -> crate::BuiltinResult<IntersectEvaluation> {
684 if a_shape.len() != 2 || b_shape.len() != 2 {
685 return Err(intersect_internal_error(
686 "intersect: 'rows' option requires 2-D numeric matrices",
687 ));
688 }
689 if a_shape[1] != b_shape[1] {
690 return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
691 }
692 let (rows_a, rows_b, cols) = (a_shape[0], b_shape[0], a_shape[1]);
693 let a_values = a_storage.exact_values();
694 let b_values = b_storage.exact_values();
695 let mut b_map = HashMap::<Vec<IntValue>, usize>::new();
696 for row in 0..rows_b {
697 let key: Vec<_> = (0..cols)
698 .map(|col| b_values[row + col * rows_b].clone())
699 .collect();
700 b_map.entry(key).or_insert(row);
701 }
702 let mut seen = HashSet::<Vec<IntValue>>::new();
703 let mut entries = Vec::<IntegerRowIntersectEntry>::new();
704 for row in 0..rows_a {
705 let values: Vec<_> = (0..cols)
706 .map(|col| a_values[row + col * rows_a].clone())
707 .collect();
708 if seen.contains(&values) {
709 continue;
710 }
711 if let Some(&b_row) = b_map.get(&values) {
712 let order_rank = entries.len();
713 entries.push(IntegerRowIntersectEntry {
714 row_data: values.clone(),
715 a_row: row,
716 b_row,
717 order_rank,
718 });
719 seen.insert(values);
720 }
721 }
722 assemble_integer_row_intersect(entries, a_storage, opts, cols)
723}
724
725fn intersect_floating_elements<T: SetFloat>(
726 a_values: Vec<T>,
727 b_values: Vec<T>,
728 opts: &IntersectOptions,
729) -> crate::BuiltinResult<IntersectEvaluation> {
730 let mut b_map: HashMap<u64, usize> = HashMap::new();
731 for (idx, &value) in b_values.iter().enumerate() {
732 let key = value.canonical_key();
733 b_map.entry(key).or_insert(idx);
734 }
735
736 let mut seen: HashSet<u64> = HashSet::new();
737 let mut entries = Vec::<FloatingIntersectEntry<T>>::new();
738 let mut order_counter = 0usize;
739
740 for (idx, &value) in a_values.iter().enumerate() {
741 let key = value.canonical_key();
742 if seen.contains(&key) {
743 continue;
744 }
745 if let Some(&b_idx) = b_map.get(&key) {
746 entries.push(FloatingIntersectEntry {
747 value,
748 a_index: idx,
749 b_index: b_idx,
750 order_rank: order_counter,
751 });
752 seen.insert(key);
753 order_counter += 1;
754 }
755 }
756
757 assemble_floating_intersect(entries, opts)
758}
759
760fn intersect_floating_rows<T: SetFloat>(
761 a_values: Vec<T>,
762 a_shape: Vec<usize>,
763 b_values: Vec<T>,
764 b_shape: Vec<usize>,
765 opts: &IntersectOptions,
766) -> crate::BuiltinResult<IntersectEvaluation> {
767 if a_shape.len() != 2 || b_shape.len() != 2 {
768 return Err(intersect_internal_error(
769 "intersect: 'rows' option requires 2-D numeric matrices",
770 ));
771 }
772 if a_shape[1] != b_shape[1] {
773 return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
774 }
775 let rows_a = a_shape[0];
776 let cols = a_shape[1];
777 let rows_b = b_shape[0];
778
779 let mut b_map: HashMap<FloatingRowKey, usize> = HashMap::new();
780 for r in 0..rows_b {
781 let mut row_values = Vec::with_capacity(cols);
782 for c in 0..cols {
783 let idx = r + c * rows_b;
784 row_values.push(b_values[idx]);
785 }
786 let key = FloatingRowKey::from_slice(&row_values);
787 b_map.entry(key).or_insert(r);
788 }
789
790 let mut seen: HashSet<FloatingRowKey> = HashSet::new();
791 let mut entries = Vec::<FloatingRowIntersectEntry<T>>::new();
792 let mut order_counter = 0usize;
793
794 for r in 0..rows_a {
795 let mut row_values = Vec::with_capacity(cols);
796 for c in 0..cols {
797 let idx = r + c * rows_a;
798 row_values.push(a_values[idx]);
799 }
800 let key = FloatingRowKey::from_slice(&row_values);
801 if seen.contains(&key) {
802 continue;
803 }
804 if let Some(&b_row) = b_map.get(&key) {
805 entries.push(FloatingRowIntersectEntry {
806 row_data: row_values,
807 a_row: r,
808 b_row,
809 order_rank: order_counter,
810 });
811 seen.insert(key);
812 order_counter += 1;
813 }
814 }
815
816 assemble_floating_row_intersect(entries, opts, cols)
817}
818
819#[cfg(test)]
820fn intersect_numeric_elements(
821 a: Tensor,
822 b: Tensor,
823 opts: &IntersectOptions,
824) -> crate::BuiltinResult<IntersectEvaluation> {
825 intersect_numeric(a, b, opts)
826}
827
828#[cfg(test)]
829fn intersect_numeric_rows(
830 a: Tensor,
831 b: Tensor,
832 opts: &IntersectOptions,
833) -> crate::BuiltinResult<IntersectEvaluation> {
834 intersect_numeric(a, b, opts)
835}
836
837fn intersect_complex(
838 a: ComplexTensor,
839 b: ComplexTensor,
840 opts: &IntersectOptions,
841) -> crate::BuiltinResult<IntersectEvaluation> {
842 let a_shape = a.shape.clone();
843 let b_shape = b.shape.clone();
844 match (a.into_complex_storage(), b.into_complex_storage()) {
845 (ComplexStorage::F64(a), ComplexStorage::F64(b)) => {
846 intersect_floating_complex(a, a_shape, b, b_shape, opts)
847 }
848 (ComplexStorage::F32(a), ComplexStorage::F32(b)) => {
849 intersect_floating_complex(a, a_shape, b, b_shape, opts)
850 }
851 (a, b) => intersect_promoted_complex_f64(a, a_shape, b, b_shape, opts),
852 }
853}
854
855fn intersect_promoted_complex_f64(
856 a: ComplexStorage,
857 a_shape: Vec<usize>,
858 b: ComplexStorage,
859 b_shape: Vec<usize>,
860 opts: &IntersectOptions,
861) -> crate::BuiltinResult<IntersectEvaluation> {
862 intersect_floating_complex(
863 a.materialize_f64(),
864 a_shape,
865 b.materialize_f64(),
866 b_shape,
867 opts,
868 )
869}
870
871fn intersect_floating_complex<T: SetFloat>(
872 a: Vec<(T, T)>,
873 a_shape: Vec<usize>,
874 b: Vec<(T, T)>,
875 b_shape: Vec<usize>,
876 opts: &IntersectOptions,
877) -> crate::BuiltinResult<IntersectEvaluation> {
878 if opts.rows {
879 intersect_complex_rows(a, a_shape, b, b_shape, opts)
880 } else {
881 intersect_complex_elements(a, b, opts)
882 }
883}
884
885fn intersect_complex_elements<T: SetFloat>(
886 a: Vec<(T, T)>,
887 b: Vec<(T, T)>,
888 opts: &IntersectOptions,
889) -> crate::BuiltinResult<IntersectEvaluation> {
890 let mut b_map: HashMap<ComplexKey, usize> = HashMap::new();
891 for (idx, &value) in b.iter().enumerate() {
892 let key = ComplexKey::new(value);
893 b_map.entry(key).or_insert(idx);
894 }
895
896 let mut seen: HashSet<ComplexKey> = HashSet::new();
897 let mut entries = Vec::<ComplexIntersectEntry<T>>::new();
898 let mut order_counter = 0usize;
899
900 for (idx, &value) in a.iter().enumerate() {
901 let key = ComplexKey::new(value);
902 if seen.contains(&key) {
903 continue;
904 }
905 if let Some(&b_idx) = b_map.get(&key) {
906 entries.push(ComplexIntersectEntry {
907 value,
908 a_index: idx,
909 b_index: b_idx,
910 order_rank: order_counter,
911 });
912 seen.insert(key);
913 order_counter += 1;
914 }
915 }
916
917 assemble_complex_intersect(entries, opts)
918}
919
920fn intersect_complex_rows<T: SetFloat>(
921 a: Vec<(T, T)>,
922 a_shape: Vec<usize>,
923 b: Vec<(T, T)>,
924 b_shape: Vec<usize>,
925 opts: &IntersectOptions,
926) -> crate::BuiltinResult<IntersectEvaluation> {
927 if a_shape.len() != 2 || b_shape.len() != 2 {
928 return Err(intersect_internal_error(
929 "intersect: 'rows' option requires 2-D complex matrices",
930 ));
931 }
932 if a_shape[1] != b_shape[1] {
933 return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
934 }
935 let rows_a = a_shape[0];
936 let cols = a_shape[1];
937 let rows_b = b_shape[0];
938
939 let mut b_map: HashMap<Vec<ComplexKey>, usize> = HashMap::new();
940 for r in 0..rows_b {
941 let mut row_keys = Vec::with_capacity(cols);
942 for c in 0..cols {
943 let idx = r + c * rows_b;
944 row_keys.push(ComplexKey::new(b[idx]));
945 }
946 b_map.entry(row_keys).or_insert(r);
947 }
948
949 let mut seen: HashSet<Vec<ComplexKey>> = HashSet::new();
950 let mut entries = Vec::<ComplexRowIntersectEntry<T>>::new();
951 let mut order_counter = 0usize;
952
953 for r in 0..rows_a {
954 let mut row_values = Vec::with_capacity(cols);
955 let mut row_keys = Vec::with_capacity(cols);
956 for c in 0..cols {
957 let idx = r + c * rows_a;
958 let value = a[idx];
959 row_values.push(value);
960 row_keys.push(ComplexKey::new(value));
961 }
962 if seen.contains(&row_keys) {
963 continue;
964 }
965 if let Some(&b_row) = b_map.get(&row_keys) {
966 entries.push(ComplexRowIntersectEntry {
967 row_data: row_values,
968 a_row: r,
969 b_row,
970 order_rank: order_counter,
971 });
972 seen.insert(row_keys);
973 order_counter += 1;
974 }
975 }
976
977 assemble_complex_row_intersect(entries, opts, cols)
978}
979
980fn intersect_char(
981 a: CharArray,
982 b: CharArray,
983 opts: &IntersectOptions,
984) -> crate::BuiltinResult<IntersectEvaluation> {
985 if opts.rows {
986 intersect_char_rows(a, b, opts)
987 } else {
988 intersect_char_elements(a, b, opts)
989 }
990}
991
992fn intersect_char_elements(
993 a: CharArray,
994 b: CharArray,
995 opts: &IntersectOptions,
996) -> crate::BuiltinResult<IntersectEvaluation> {
997 let mut seen: HashSet<u32> = HashSet::new();
998 let mut entries = Vec::<CharIntersectEntry>::new();
999 let mut order_counter = 0usize;
1000
1001 for col in 0..a.cols {
1002 for row in 0..a.rows {
1003 let linear_idx = row + col * a.rows;
1004 let data_idx = row * a.cols + col;
1005 let ch = a.data[data_idx];
1006 let key = ch as u32;
1007 if seen.contains(&key) {
1008 continue;
1009 }
1010 if let Some(b_idx) = find_char_index(&b, ch) {
1011 entries.push(CharIntersectEntry {
1012 ch,
1013 a_index: linear_idx,
1014 b_index: b_idx,
1015 order_rank: order_counter,
1016 });
1017 seen.insert(key);
1018 order_counter += 1;
1019 }
1020 }
1021 }
1022
1023 assemble_char_intersect(entries, opts, &b)
1024}
1025
1026fn intersect_char_rows(
1027 a: CharArray,
1028 b: CharArray,
1029 opts: &IntersectOptions,
1030) -> crate::BuiltinResult<IntersectEvaluation> {
1031 if a.cols != b.cols {
1032 return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
1033 }
1034 let rows_a = a.rows;
1035 let rows_b = b.rows;
1036 let cols = a.cols;
1037
1038 let mut b_map: HashMap<RowCharKey, usize> = HashMap::new();
1039 for r in 0..rows_b {
1040 let mut row_values = Vec::with_capacity(cols);
1041 for c in 0..cols {
1042 let idx = r * cols + c;
1043 row_values.push(b.data[idx]);
1044 }
1045 let key = RowCharKey::from_slice(&row_values);
1046 b_map.entry(key).or_insert(r);
1047 }
1048
1049 let mut seen: HashSet<RowCharKey> = HashSet::new();
1050 let mut entries = Vec::<CharRowIntersectEntry>::new();
1051 let mut order_counter = 0usize;
1052
1053 for r in 0..rows_a {
1054 let mut row_values = Vec::with_capacity(cols);
1055 for c in 0..cols {
1056 let idx = r * cols + c;
1057 row_values.push(a.data[idx]);
1058 }
1059 let key = RowCharKey::from_slice(&row_values);
1060 if seen.contains(&key) {
1061 continue;
1062 }
1063 if let Some(&b_row) = b_map.get(&key) {
1064 entries.push(CharRowIntersectEntry {
1065 row_data: row_values,
1066 a_row: r,
1067 b_row,
1068 order_rank: order_counter,
1069 });
1070 seen.insert(key);
1071 order_counter += 1;
1072 }
1073 }
1074
1075 assemble_char_row_intersect(entries, opts, cols)
1076}
1077
1078fn find_char_index(array: &CharArray, target: char) -> Option<usize> {
1079 for col in 0..array.cols {
1080 for row in 0..array.rows {
1081 let data_idx = row * array.cols + col;
1082 if array.data[data_idx] == target {
1083 return Some(row + col * array.rows);
1084 }
1085 }
1086 }
1087 None
1088}
1089
1090fn intersect_string(
1091 a: StringArray,
1092 b: StringArray,
1093 opts: &IntersectOptions,
1094) -> crate::BuiltinResult<IntersectEvaluation> {
1095 if opts.rows {
1096 intersect_string_rows(a, b, opts)
1097 } else {
1098 intersect_string_elements(a, b, opts)
1099 }
1100}
1101
1102fn intersect_string_elements(
1103 a: StringArray,
1104 b: StringArray,
1105 opts: &IntersectOptions,
1106) -> crate::BuiltinResult<IntersectEvaluation> {
1107 let mut b_map: HashMap<String, usize> = HashMap::new();
1108 for (idx, value) in b.data.iter().enumerate() {
1109 b_map.entry(value.clone()).or_insert(idx);
1110 }
1111
1112 let mut seen: HashSet<String> = HashSet::new();
1113 let mut entries = Vec::<StringIntersectEntry>::new();
1114 let mut order_counter = 0usize;
1115
1116 for (idx, value) in a.data.iter().enumerate() {
1117 if seen.contains(value) {
1118 continue;
1119 }
1120 if let Some(&b_idx) = b_map.get(value) {
1121 entries.push(StringIntersectEntry {
1122 value: value.clone(),
1123 a_index: idx,
1124 b_index: b_idx,
1125 order_rank: order_counter,
1126 });
1127 seen.insert(value.clone());
1128 order_counter += 1;
1129 }
1130 }
1131
1132 assemble_string_intersect(entries, opts)
1133}
1134
1135fn intersect_string_rows(
1136 a: StringArray,
1137 b: StringArray,
1138 opts: &IntersectOptions,
1139) -> crate::BuiltinResult<IntersectEvaluation> {
1140 if a.shape.len() != 2 || b.shape.len() != 2 {
1141 return Err(intersect_internal_error(
1142 "intersect: 'rows' option requires 2-D string arrays",
1143 ));
1144 }
1145 if a.shape[1] != b.shape[1] {
1146 return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
1147 }
1148 let rows_a = a.shape[0];
1149 let cols = a.shape[1];
1150 let rows_b = b.shape[0];
1151
1152 let mut b_map: HashMap<RowStringKey, usize> = HashMap::new();
1153 for r in 0..rows_b {
1154 let mut row_values = Vec::with_capacity(cols);
1155 for c in 0..cols {
1156 let idx = r + c * rows_b;
1157 row_values.push(b.data[idx].clone());
1158 }
1159 let key = RowStringKey::from_slice(&row_values);
1160 b_map.entry(key).or_insert(r);
1161 }
1162
1163 let mut seen: HashSet<RowStringKey> = HashSet::new();
1164 let mut entries = Vec::<StringRowIntersectEntry>::new();
1165 let mut order_counter = 0usize;
1166
1167 for r in 0..rows_a {
1168 let mut row_values = Vec::with_capacity(cols);
1169 for c in 0..cols {
1170 let idx = r + c * rows_a;
1171 row_values.push(a.data[idx].clone());
1172 }
1173 let key = RowStringKey::from_slice(&row_values);
1174 if seen.contains(&key) {
1175 continue;
1176 }
1177 if let Some(&b_row) = b_map.get(&key) {
1178 entries.push(StringRowIntersectEntry {
1179 row_data: row_values,
1180 a_row: r,
1181 b_row,
1182 order_rank: order_counter,
1183 });
1184 seen.insert(key);
1185 order_counter += 1;
1186 }
1187 }
1188
1189 assemble_string_row_intersect(entries, opts, cols)
1190}
1191
1192#[derive(Debug, Clone)]
1193pub struct IntersectEvaluation {
1194 values: Value,
1195 ia: Tensor,
1196 ib: Tensor,
1197}
1198
1199impl IntersectEvaluation {
1200 fn new(values: Value, ia: Tensor, ib: Tensor) -> Self {
1201 Self { values, ia, ib }
1202 }
1203
1204 pub fn into_values_value(self) -> Value {
1205 self.values
1206 }
1207
1208 pub fn into_pair(self) -> (Value, Value) {
1209 let ia = tensor::tensor_into_value(self.ia);
1210 (self.values, ia)
1211 }
1212
1213 pub fn into_triple(self) -> (Value, Value, Value) {
1214 let ia = tensor::tensor_into_value(self.ia);
1215 let ib = tensor::tensor_into_value(self.ib);
1216 (self.values, ia, ib)
1217 }
1218
1219 pub fn values_value(&self) -> Value {
1220 self.values.clone()
1221 }
1222
1223 pub fn ia_value(&self) -> Value {
1224 tensor::tensor_into_value(self.ia.clone())
1225 }
1226
1227 pub fn ib_value(&self) -> Value {
1228 tensor::tensor_into_value(self.ib.clone())
1229 }
1230}
1231
1232#[derive(Debug)]
1233struct FloatingIntersectEntry<T> {
1234 value: T,
1235 a_index: usize,
1236 b_index: usize,
1237 order_rank: usize,
1238}
1239
1240#[derive(Debug)]
1241struct IntegerIntersectEntry {
1242 value: IntValue,
1243 a_index: usize,
1244 b_index: usize,
1245 order_rank: usize,
1246}
1247
1248#[derive(Debug)]
1249struct FloatingRowIntersectEntry<T> {
1250 row_data: Vec<T>,
1251 a_row: usize,
1252 b_row: usize,
1253 order_rank: usize,
1254}
1255
1256#[derive(Debug)]
1257struct IntegerRowIntersectEntry {
1258 row_data: Vec<IntValue>,
1259 a_row: usize,
1260 b_row: usize,
1261 order_rank: usize,
1262}
1263
1264#[derive(Debug)]
1265struct ComplexIntersectEntry<T> {
1266 value: (T, T),
1267 a_index: usize,
1268 b_index: usize,
1269 order_rank: usize,
1270}
1271
1272#[derive(Debug)]
1273struct ComplexRowIntersectEntry<T> {
1274 row_data: Vec<(T, T)>,
1275 a_row: usize,
1276 b_row: usize,
1277 order_rank: usize,
1278}
1279
1280#[derive(Debug)]
1281struct CharIntersectEntry {
1282 ch: char,
1283 a_index: usize,
1284 b_index: usize,
1285 order_rank: usize,
1286}
1287
1288#[derive(Debug)]
1289struct CharRowIntersectEntry {
1290 row_data: Vec<char>,
1291 a_row: usize,
1292 b_row: usize,
1293 order_rank: usize,
1294}
1295
1296#[derive(Debug)]
1297struct StringIntersectEntry {
1298 value: String,
1299 a_index: usize,
1300 b_index: usize,
1301 order_rank: usize,
1302}
1303
1304#[derive(Debug)]
1305struct StringRowIntersectEntry {
1306 row_data: Vec<String>,
1307 a_row: usize,
1308 b_row: usize,
1309 order_rank: usize,
1310}
1311
1312fn assemble_floating_intersect<T: SetFloat>(
1313 entries: Vec<FloatingIntersectEntry<T>>,
1314 opts: &IntersectOptions,
1315) -> crate::BuiltinResult<IntersectEvaluation> {
1316 let mut order: Vec<usize> = (0..entries.len()).collect();
1317 match opts.order {
1318 IntersectOrder::Sorted => {
1319 order.sort_by(|&lhs, &rhs| entries[lhs].value.compare(entries[rhs].value));
1320 }
1321 IntersectOrder::Stable => {
1322 order.sort_by_key(|&idx| entries[idx].order_rank);
1323 }
1324 }
1325
1326 let mut values = Vec::with_capacity(order.len());
1327 let mut ia = Vec::with_capacity(order.len());
1328 let mut ib = Vec::with_capacity(order.len());
1329 for &idx in &order {
1330 let entry = &entries[idx];
1331 values.push(entry.value);
1332 ia.push((entry.a_index + 1) as f64);
1333 ib.push((entry.b_index + 1) as f64);
1334 }
1335
1336 let value_tensor =
1337 Tensor::from_numeric_storage(T::numeric_storage(values), vec![order.len(), 1])
1338 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1339 let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1340 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1341 let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1342 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1343
1344 Ok(IntersectEvaluation::new(
1345 tensor::tensor_into_value(value_tensor),
1346 ia_tensor,
1347 ib_tensor,
1348 ))
1349}
1350
1351fn assemble_integer_intersect(
1352 entries: Vec<IntegerIntersectEntry>,
1353 storage: &IntegerStorage,
1354 opts: &IntersectOptions,
1355) -> crate::BuiltinResult<IntersectEvaluation> {
1356 let mut order: Vec<_> = (0..entries.len()).collect();
1357 match opts.order {
1358 IntersectOrder::Sorted => order.sort_by(|&a, &b| {
1359 integer_order::compare(&entries[a].value, &entries[b].value, false, false)
1360 }),
1361 IntersectOrder::Stable => order.sort_by_key(|&index| entries[index].order_rank),
1362 }
1363 let values: Vec<_> = order
1364 .iter()
1365 .map(|&index| entries[index].value.clone())
1366 .collect();
1367 let ia: Vec<_> = order
1368 .iter()
1369 .map(|&index| (entries[index].a_index + 1) as f64)
1370 .collect();
1371 let ib: Vec<_> = order
1372 .iter()
1373 .map(|&index| (entries[index].b_index + 1) as f64)
1374 .collect();
1375 let values = Tensor::new_integer(
1376 storage
1377 .from_exact_values_like(values)
1378 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?,
1379 vec![order.len(), 1],
1380 )
1381 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1382 let ia = Tensor::new(ia, vec![order.len(), 1])
1383 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1384 let ib = Tensor::new(ib, vec![order.len(), 1])
1385 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1386 Ok(IntersectEvaluation::new(Value::Tensor(values), ia, ib))
1387}
1388
1389fn assemble_floating_row_intersect<T: SetFloat>(
1390 entries: Vec<FloatingRowIntersectEntry<T>>,
1391 opts: &IntersectOptions,
1392 cols: usize,
1393) -> crate::BuiltinResult<IntersectEvaluation> {
1394 let mut order: Vec<usize> = (0..entries.len()).collect();
1395 match opts.order {
1396 IntersectOrder::Sorted => {
1397 order.sort_by(|&lhs, &rhs| {
1398 compare_floating_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1399 });
1400 }
1401 IntersectOrder::Stable => {
1402 order.sort_by_key(|&idx| entries[idx].order_rank);
1403 }
1404 }
1405
1406 let rows_out = order.len();
1407 let mut values = vec![T::default(); rows_out * cols];
1408 let mut ia = Vec::with_capacity(rows_out);
1409 let mut ib = Vec::with_capacity(rows_out);
1410
1411 for (row_pos, &entry_idx) in order.iter().enumerate() {
1412 let entry = &entries[entry_idx];
1413 for col in 0..cols {
1414 let dest = row_pos + col * rows_out;
1415 values[dest] = entry.row_data[col];
1416 }
1417 ia.push((entry.a_row + 1) as f64);
1418 ib.push((entry.b_row + 1) as f64);
1419 }
1420
1421 let value_tensor =
1422 Tensor::from_numeric_storage(T::numeric_storage(values), vec![rows_out, cols])
1423 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1424 let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1425 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1426 let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1427 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1428
1429 Ok(IntersectEvaluation::new(
1430 tensor::tensor_into_value(value_tensor),
1431 ia_tensor,
1432 ib_tensor,
1433 ))
1434}
1435
1436fn assemble_integer_row_intersect(
1437 entries: Vec<IntegerRowIntersectEntry>,
1438 storage: &IntegerStorage,
1439 opts: &IntersectOptions,
1440 cols: usize,
1441) -> crate::BuiltinResult<IntersectEvaluation> {
1442 let mut order: Vec<_> = (0..entries.len()).collect();
1443 match opts.order {
1444 IntersectOrder::Sorted => order.sort_by(|&a, &b| {
1445 for (left, right) in entries[a].row_data.iter().zip(&entries[b].row_data) {
1446 let ordering = integer_order::compare(left, right, false, false);
1447 if ordering != Ordering::Equal {
1448 return ordering;
1449 }
1450 }
1451 Ordering::Equal
1452 }),
1453 IntersectOrder::Stable => order.sort_by_key(|&index| entries[index].order_rank),
1454 }
1455 let rows = order.len();
1456 let mut values = Vec::with_capacity(rows * cols);
1457 for col in 0..cols {
1458 for &index in &order {
1459 values.push(entries[index].row_data[col].clone());
1460 }
1461 }
1462 let ia: Vec<_> = order
1463 .iter()
1464 .map(|&index| (entries[index].a_row + 1) as f64)
1465 .collect();
1466 let ib: Vec<_> = order
1467 .iter()
1468 .map(|&index| (entries[index].b_row + 1) as f64)
1469 .collect();
1470 let values = Tensor::new_integer(
1471 storage
1472 .from_exact_values_like(values)
1473 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?,
1474 vec![rows, cols],
1475 )
1476 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1477 let ia = Tensor::new(ia, vec![rows, 1])
1478 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1479 let ib = Tensor::new(ib, vec![rows, 1])
1480 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1481 Ok(IntersectEvaluation::new(Value::Tensor(values), ia, ib))
1482}
1483
1484fn assemble_complex_intersect<T: SetFloat>(
1485 entries: Vec<ComplexIntersectEntry<T>>,
1486 opts: &IntersectOptions,
1487) -> crate::BuiltinResult<IntersectEvaluation> {
1488 let mut order: Vec<usize> = (0..entries.len()).collect();
1489 match opts.order {
1490 IntersectOrder::Sorted => {
1491 order.sort_by(|&lhs, &rhs| compare_complex(entries[lhs].value, entries[rhs].value));
1492 }
1493 IntersectOrder::Stable => {
1494 order.sort_by_key(|&idx| entries[idx].order_rank);
1495 }
1496 }
1497
1498 let mut values = Vec::with_capacity(order.len());
1499 let mut ia = Vec::with_capacity(order.len());
1500 let mut ib = Vec::with_capacity(order.len());
1501 for &idx in &order {
1502 let entry = &entries[idx];
1503 values.push(entry.value);
1504 ia.push((entry.a_index + 1) as f64);
1505 ib.push((entry.b_index + 1) as f64);
1506 }
1507
1508 let value_tensor =
1509 ComplexTensor::from_complex_storage(T::complex_storage(values), vec![order.len(), 1])
1510 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1511 let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1512 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1513 let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1514 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1515
1516 let value = if value_tensor.as_f32_slice().is_some() {
1517 Value::ComplexTensor(value_tensor)
1518 } else {
1519 complex_tensor_into_value(value_tensor)
1520 };
1521 Ok(IntersectEvaluation::new(value, ia_tensor, ib_tensor))
1522}
1523
1524fn assemble_complex_row_intersect<T: SetFloat>(
1525 entries: Vec<ComplexRowIntersectEntry<T>>,
1526 opts: &IntersectOptions,
1527 cols: usize,
1528) -> crate::BuiltinResult<IntersectEvaluation> {
1529 let mut order: Vec<usize> = (0..entries.len()).collect();
1530 match opts.order {
1531 IntersectOrder::Sorted => {
1532 order.sort_by(|&lhs, &rhs| {
1533 compare_complex_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1534 });
1535 }
1536 IntersectOrder::Stable => {
1537 order.sort_by_key(|&idx| entries[idx].order_rank);
1538 }
1539 }
1540
1541 let rows_out = order.len();
1542 let mut values = vec![(T::default(), T::default()); rows_out * cols];
1543 let mut ia = Vec::with_capacity(rows_out);
1544 let mut ib = Vec::with_capacity(rows_out);
1545
1546 for (row_pos, &entry_idx) in order.iter().enumerate() {
1547 let entry = &entries[entry_idx];
1548 for col in 0..cols {
1549 let dest = row_pos + col * rows_out;
1550 values[dest] = entry.row_data[col];
1551 }
1552 ia.push((entry.a_row + 1) as f64);
1553 ib.push((entry.b_row + 1) as f64);
1554 }
1555
1556 let value_tensor =
1557 ComplexTensor::from_complex_storage(T::complex_storage(values), vec![rows_out, cols])
1558 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1559 let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1560 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1561 let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1562 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1563
1564 let value = if value_tensor.as_f32_slice().is_some() {
1565 Value::ComplexTensor(value_tensor)
1566 } else {
1567 complex_tensor_into_value(value_tensor)
1568 };
1569 Ok(IntersectEvaluation::new(value, ia_tensor, ib_tensor))
1570}
1571
1572fn assemble_char_intersect(
1573 entries: Vec<CharIntersectEntry>,
1574 opts: &IntersectOptions,
1575 b: &CharArray,
1576) -> crate::BuiltinResult<IntersectEvaluation> {
1577 let mut order: Vec<usize> = (0..entries.len()).collect();
1578 match opts.order {
1579 IntersectOrder::Sorted => {
1580 order.sort_by(|&lhs, &rhs| entries[lhs].ch.cmp(&entries[rhs].ch));
1581 }
1582 IntersectOrder::Stable => {
1583 order.sort_by_key(|&idx| entries[idx].order_rank);
1584 }
1585 }
1586
1587 let mut values = Vec::with_capacity(order.len());
1588 let mut ia = Vec::with_capacity(order.len());
1589 let mut ib = Vec::with_capacity(order.len());
1590 for &idx in &order {
1591 let entry = &entries[idx];
1592 values.push(entry.ch);
1593 ia.push((entry.a_index + 1) as f64);
1594 let b_idx = find_char_index(b, entry.ch).unwrap_or(entry.b_index);
1595 ib.push((b_idx + 1) as f64);
1596 }
1597
1598 let value_array = CharArray::new(values, order.len(), 1)
1599 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1600 let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1601 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1602 let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1603 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1604
1605 Ok(IntersectEvaluation::new(
1606 Value::CharArray(value_array),
1607 ia_tensor,
1608 ib_tensor,
1609 ))
1610}
1611
1612fn assemble_char_row_intersect(
1613 entries: Vec<CharRowIntersectEntry>,
1614 opts: &IntersectOptions,
1615 cols: usize,
1616) -> crate::BuiltinResult<IntersectEvaluation> {
1617 let mut order: Vec<usize> = (0..entries.len()).collect();
1618 match opts.order {
1619 IntersectOrder::Sorted => {
1620 order.sort_by(|&lhs, &rhs| {
1621 compare_char_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1622 });
1623 }
1624 IntersectOrder::Stable => {
1625 order.sort_by_key(|&idx| entries[idx].order_rank);
1626 }
1627 }
1628
1629 let rows_out = order.len();
1630 let mut values = vec!['\0'; rows_out * cols];
1631 let mut ia = Vec::with_capacity(rows_out);
1632 let mut ib = Vec::with_capacity(rows_out);
1633
1634 for (row_pos, &entry_idx) in order.iter().enumerate() {
1635 let entry = &entries[entry_idx];
1636 for col in 0..cols {
1637 let dest = row_pos * cols + col;
1638 values[dest] = entry.row_data[col];
1639 }
1640 ia.push((entry.a_row + 1) as f64);
1641 ib.push((entry.b_row + 1) as f64);
1642 }
1643
1644 let value_array = CharArray::new(values, rows_out, cols)
1645 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1646 let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1647 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1648 let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1649 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1650
1651 Ok(IntersectEvaluation::new(
1652 Value::CharArray(value_array),
1653 ia_tensor,
1654 ib_tensor,
1655 ))
1656}
1657
1658fn assemble_string_intersect(
1659 entries: Vec<StringIntersectEntry>,
1660 opts: &IntersectOptions,
1661) -> crate::BuiltinResult<IntersectEvaluation> {
1662 let mut order: Vec<usize> = (0..entries.len()).collect();
1663 match opts.order {
1664 IntersectOrder::Sorted => {
1665 order.sort_by(|&lhs, &rhs| entries[lhs].value.cmp(&entries[rhs].value));
1666 }
1667 IntersectOrder::Stable => {
1668 order.sort_by_key(|&idx| entries[idx].order_rank);
1669 }
1670 }
1671
1672 let mut values = Vec::with_capacity(order.len());
1673 let mut ia = Vec::with_capacity(order.len());
1674 let mut ib = Vec::with_capacity(order.len());
1675 for &idx in &order {
1676 let entry = &entries[idx];
1677 values.push(entry.value.clone());
1678 ia.push((entry.a_index + 1) as f64);
1679 ib.push((entry.b_index + 1) as f64);
1680 }
1681
1682 let value_array = StringArray::new(values, vec![order.len(), 1])
1683 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1684 let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1685 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1686 let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1687 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1688
1689 Ok(IntersectEvaluation::new(
1690 Value::StringArray(value_array),
1691 ia_tensor,
1692 ib_tensor,
1693 ))
1694}
1695
1696fn assemble_string_row_intersect(
1697 entries: Vec<StringRowIntersectEntry>,
1698 opts: &IntersectOptions,
1699 cols: usize,
1700) -> crate::BuiltinResult<IntersectEvaluation> {
1701 let mut order: Vec<usize> = (0..entries.len()).collect();
1702 match opts.order {
1703 IntersectOrder::Sorted => {
1704 order.sort_by(|&lhs, &rhs| {
1705 compare_string_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1706 });
1707 }
1708 IntersectOrder::Stable => {
1709 order.sort_by_key(|&idx| entries[idx].order_rank);
1710 }
1711 }
1712
1713 let rows_out = order.len();
1714 let mut values = vec![String::new(); rows_out * cols];
1715 let mut ia = Vec::with_capacity(rows_out);
1716 let mut ib = Vec::with_capacity(rows_out);
1717
1718 for (row_pos, &entry_idx) in order.iter().enumerate() {
1719 let entry = &entries[entry_idx];
1720 for col in 0..cols {
1721 let dest = row_pos + col * rows_out;
1722 values[dest] = entry.row_data[col].clone();
1723 }
1724 ia.push((entry.a_row + 1) as f64);
1725 ib.push((entry.b_row + 1) as f64);
1726 }
1727
1728 let value_array = StringArray::new(values, vec![rows_out, cols])
1729 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1730 let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1731 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1732 let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1733 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1734
1735 Ok(IntersectEvaluation::new(
1736 Value::StringArray(value_array),
1737 ia_tensor,
1738 ib_tensor,
1739 ))
1740}
1741
1742#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1743struct FloatingRowKey(Vec<u64>);
1744
1745impl FloatingRowKey {
1746 fn from_slice<T: SetFloat>(values: &[T]) -> Self {
1747 Self(values.iter().map(|&value| value.canonical_key()).collect())
1748 }
1749}
1750
1751#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
1752struct ComplexKey {
1753 re: u64,
1754 im: u64,
1755}
1756
1757impl ComplexKey {
1758 fn new<T: SetFloat>(value: (T, T)) -> Self {
1759 Self {
1760 re: value.0.canonical_key(),
1761 im: value.1.canonical_key(),
1762 }
1763 }
1764}
1765
1766#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1767struct RowCharKey(Vec<u32>);
1768
1769impl RowCharKey {
1770 fn from_slice(values: &[char]) -> Self {
1771 RowCharKey(values.iter().map(|&ch| ch as u32).collect())
1772 }
1773}
1774
1775#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1776struct RowStringKey(Vec<String>);
1777
1778impl RowStringKey {
1779 fn from_slice(values: &[String]) -> Self {
1780 RowStringKey(values.to_vec())
1781 }
1782}
1783
1784fn scalar_complex_tensor(re: f64, im: f64) -> crate::BuiltinResult<ComplexTensor> {
1785 ComplexTensor::new(vec![(re, im)], vec![1, 1])
1786 .map_err(|e| intersect_internal_error(format!("intersect: {e}")))
1787}
1788
1789fn tensor_to_complex_owned(name: &str, tensor: Tensor) -> crate::BuiltinResult<ComplexTensor> {
1790 let shape = tensor.shape.clone();
1791 let complex = tensor
1792 .into_numeric_storage()
1793 .map_err(|e| intersect_internal_error(format!("{name}: {e}")))?
1794 .materialize_f64()
1795 .into_iter()
1796 .map(|real| (real, 0.0))
1797 .collect();
1798 ComplexTensor::new(complex, shape).map_err(|e| intersect_internal_error(format!("{name}: {e}")))
1799}
1800
1801fn value_into_complex_tensor(value: Value) -> crate::BuiltinResult<ComplexTensor> {
1802 match value {
1803 Value::ComplexTensor(tensor) => Ok(tensor),
1804 Value::Complex(re, im) => scalar_complex_tensor(re, im),
1805 other => {
1806 let tensor = tensor::value_into_tensor_for("intersect", other)
1807 .map_err(|e| intersect_internal_error(e))?;
1808 tensor_to_complex_owned("intersect", tensor)
1809 }
1810 }
1811}
1812
1813fn compare_floating_rows<T: SetFloat>(a: &[T], b: &[T]) -> Ordering {
1814 for (lhs, rhs) in a.iter().zip(b.iter()) {
1815 let ord = lhs.compare(*rhs);
1816 if ord != Ordering::Equal {
1817 return ord;
1818 }
1819 }
1820 Ordering::Equal
1821}
1822
1823fn complex_is_nan<T: SetFloat>(value: (T, T)) -> bool {
1824 value.0.is_nan() || value.1.is_nan()
1825}
1826
1827fn compare_complex<T: SetFloat>(a: (T, T), b: (T, T)) -> Ordering {
1828 match (complex_is_nan(a), complex_is_nan(b)) {
1829 (true, true) => Ordering::Equal,
1830 (true, false) => Ordering::Greater,
1831 (false, true) => Ordering::Less,
1832 (false, false) => {
1833 let mag_a = a.0.hypot(a.1);
1834 let mag_b = b.0.hypot(b.1);
1835 let mag_cmp = mag_a.compare(mag_b);
1836 if mag_cmp != Ordering::Equal {
1837 return mag_cmp;
1838 }
1839 let re_cmp = a.0.compare(b.0);
1840 if re_cmp != Ordering::Equal {
1841 return re_cmp;
1842 }
1843 a.1.compare(b.1)
1844 }
1845 }
1846}
1847
1848fn compare_complex_rows<T: SetFloat>(a: &[(T, T)], b: &[(T, T)]) -> Ordering {
1849 for (lhs, rhs) in a.iter().zip(b.iter()) {
1850 let ord = compare_complex(*lhs, *rhs);
1851 if ord != Ordering::Equal {
1852 return ord;
1853 }
1854 }
1855 Ordering::Equal
1856}
1857
1858fn compare_char_rows(a: &[char], b: &[char]) -> Ordering {
1859 for (lhs, rhs) in a.iter().zip(b.iter()) {
1860 let ord = lhs.cmp(rhs);
1861 if ord != Ordering::Equal {
1862 return ord;
1863 }
1864 }
1865 Ordering::Equal
1866}
1867
1868fn compare_string_rows(a: &[String], b: &[String]) -> Ordering {
1869 for (lhs, rhs) in a.iter().zip(b.iter()) {
1870 let ord = lhs.cmp(rhs);
1871 if ord != Ordering::Equal {
1872 return ord;
1873 }
1874 }
1875 Ordering::Equal
1876}
1877
1878#[cfg(test)]
1879pub(crate) mod tests {
1880 use super::*;
1881 use crate::builtins::common::test_support;
1882 use runmat_accelerate_api::HostTensorView;
1883 use runmat_builtins::{ResolveContext, Type};
1884
1885 fn evaluate_sync(
1886 a: Value,
1887 b: Value,
1888 rest: &[Value],
1889 ) -> crate::BuiltinResult<IntersectEvaluation> {
1890 futures::executor::block_on(evaluate(a, b, rest))
1891 }
1892
1893 fn builtin_sync(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1894 futures::executor::block_on(intersect_builtin(a, b, rest))
1895 }
1896
1897 #[test]
1898 fn registered_builtin_restores_resident_outputs_and_rejects_excess_arity() {
1899 test_support::with_test_provider(|provider| {
1900 let left = Tensor::new_integer(IntegerStorage::I32(vec![7, 2, 9]), vec![3, 1]).unwrap();
1901 let right = Tensor::new_integer(IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
1902 let left =
1903 Value::GpuTensor(gpu_helpers::upload_tensor(provider, &left).expect("upload left"));
1904 let right = Value::GpuTensor(
1905 gpu_helpers::upload_tensor(provider, &right).expect("upload right"),
1906 );
1907
1908 {
1909 let _guard = crate::output_count::push_output_count(Some(3));
1910 let Value::OutputList(outputs) =
1911 builtin_sync(left, right, Vec::new()).expect("resident intersect")
1912 else {
1913 panic!("expected output list");
1914 };
1915 assert_eq!(outputs.len(), 3);
1916 assert!(outputs
1917 .iter()
1918 .all(|output| matches!(output, Value::GpuTensor(_))));
1919 assert_eq!(
1920 test_support::gather(outputs[0].clone())
1921 .expect("gather values")
1922 .integer_storage(),
1923 Some(&IntegerStorage::I32(vec![2, 7]))
1924 );
1925 }
1926
1927 let _guard = crate::output_count::push_output_count(Some(4));
1928 let err = builtin_sync(Value::Num(1.0), Value::Num(1.0), Vec::new())
1929 .expect_err("excess outputs must fail");
1930 assert_eq!(
1931 err.identifier(),
1932 INTERSECT_ERROR_INVALID_ARGUMENT.identifier
1933 );
1934 });
1935 }
1936
1937 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1938 #[test]
1939 fn intersect_numeric_sorted() {
1940 let a = Tensor::new(vec![5.0, 7.0, 5.0, 1.0], vec![4, 1]).unwrap();
1941 let b = Tensor::new(vec![7.0, 1.0, 3.0], vec![3, 1]).unwrap();
1942 let eval = intersect_numeric_elements(
1943 a,
1944 b,
1945 &IntersectOptions {
1946 rows: false,
1947 order: IntersectOrder::Sorted,
1948 },
1949 )
1950 .expect("intersect");
1951 let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
1952 assert_eq!(values.materialize_f64(), vec![1.0, 7.0]);
1953 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
1954 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
1955 assert_eq!(ia.materialize_f64(), vec![4.0, 2.0]);
1956 assert_eq!(ib.materialize_f64(), vec![2.0, 1.0]);
1957 }
1958
1959 #[test]
1960 fn intersect_preserves_native_single_elements_and_rows() {
1961 let a = Tensor::from_f32(vec![5.0, 7.0, 5.0, 1.0], vec![4, 1]).unwrap();
1962 let b = Tensor::from_f32(vec![7.0, 1.0, 3.0], vec![3, 1]).unwrap();
1963 let values = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
1964 .expect("single intersect")
1965 .values_value();
1966 let Value::Tensor(values) = values else {
1967 panic!("expected native single values");
1968 };
1969 assert_eq!(
1970 values.into_numeric_storage().unwrap(),
1971 NumericStorage::F32(vec![1.0, 7.0])
1972 );
1973
1974 let a = Tensor::from_f32(vec![1.0, 3.0, 1.0, 2.0, 4.0, 2.0], vec![3, 2]).unwrap();
1975 let b = Tensor::from_f32(vec![1.0, 5.0, 2.0, 6.0], vec![2, 2]).unwrap();
1976 let values = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")])
1977 .expect("single row intersect")
1978 .values_value();
1979 let Value::Tensor(values) = values else {
1980 panic!("expected native single rows");
1981 };
1982 assert_eq!(values.shape, vec![1, 2]);
1983 assert_eq!(
1984 values.into_numeric_storage().unwrap(),
1985 NumericStorage::F32(vec![1.0, 2.0])
1986 );
1987 }
1988
1989 #[test]
1990 fn intersect_preserves_native_complex_single_elements_and_rows() {
1991 let a =
1992 ComplexTensor::from_f32(vec![(1.0, 1.0), (0.0, 2.0), (1.0, -1.0)], vec![3, 1]).unwrap();
1993 let b = ComplexTensor::from_f32(vec![(0.0, 2.0), (4.0, 0.0)], vec![2, 1]).unwrap();
1994 let values = evaluate_sync(Value::ComplexTensor(a), Value::ComplexTensor(b), &[])
1995 .expect("complex single intersect")
1996 .values_value();
1997 let Value::ComplexTensor(values) = values else {
1998 panic!("expected native complex single value");
1999 };
2000 assert_eq!(values.as_f32_slice(), Some(&[(0.0, 2.0)][..]));
2001
2002 let a = ComplexTensor::from_f32(
2003 vec![
2004 (1.0, 0.0),
2005 (3.0, 0.0),
2006 (1.0, 0.0),
2007 (2.0, 1.0),
2008 (4.0, 1.0),
2009 (2.0, 1.0),
2010 ],
2011 vec![3, 2],
2012 )
2013 .unwrap();
2014 let b = ComplexTensor::from_f32(
2015 vec![(1.0, 0.0), (5.0, 0.0), (2.0, 1.0), (6.0, 1.0)],
2016 vec![2, 2],
2017 )
2018 .unwrap();
2019 let values = evaluate_sync(
2020 Value::ComplexTensor(a),
2021 Value::ComplexTensor(b),
2022 &[Value::from("rows")],
2023 )
2024 .expect("complex single row intersect")
2025 .values_value();
2026 let Value::ComplexTensor(values) = values else {
2027 panic!("expected native complex single rows");
2028 };
2029 assert_eq!(values.shape, vec![1, 2]);
2030 assert_eq!(values.as_f32_slice(), Some(&[(1.0, 0.0), (2.0, 1.0)][..]));
2031 }
2032
2033 #[test]
2034 fn intersect_preserves_exact_integer_elements_and_rows() {
2035 let a = Tensor::new_integer(
2036 runmat_value::IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993]),
2037 vec![3, 1],
2038 )
2039 .expect("input");
2040 let b = Tensor::new_integer(
2041 runmat_value::IntegerStorage::U64(vec![0, u64::MAX]),
2042 vec![2, 1],
2043 )
2044 .expect("input");
2045 let (values, ia, ib) = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
2046 .expect("intersect")
2047 .into_triple();
2048 let Value::Tensor(values) = values else {
2049 panic!("exact values");
2050 };
2051 assert_eq!(
2052 values.integer_storage(),
2053 Some(&runmat_value::IntegerStorage::U64(vec![0, u64::MAX]))
2054 );
2055 let Value::Tensor(ia) = ia else {
2056 panic!("indices");
2057 };
2058 assert_eq!(ia.materialize_f64(), vec![2.0, 1.0]);
2059 let Value::Tensor(ib) = ib else {
2060 panic!("indices");
2061 };
2062 assert_eq!(ib.materialize_f64(), vec![1.0, 2.0]);
2063 }
2064
2065 #[test]
2066 fn intersect_rejects_mixed_nondouble_integer_classes() {
2067 let a = Tensor::new_integer(IntegerStorage::U16(vec![7, 2, 9, 7]), vec![4, 1]).unwrap();
2068 let b = Tensor::new_integer(IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
2069
2070 let error = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
2071 .expect_err("mixed integer classes must reject");
2072 assert_eq!(
2073 error.identifier(),
2074 INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH.identifier
2075 );
2076 }
2077
2078 #[test]
2079 fn intersect_type_resolver_numeric() {
2080 assert_eq!(
2081 set_values_output_type(&[Type::tensor()], &ResolveContext::new(Vec::new())),
2082 Type::tensor()
2083 );
2084 }
2085
2086 #[test]
2087 fn intersect_type_resolver_string_array() {
2088 assert_eq!(
2089 set_values_output_type(
2090 &[Type::cell_of(Type::String)],
2091 &ResolveContext::new(Vec::new()),
2092 ),
2093 Type::cell_of(Type::String)
2094 );
2095 }
2096
2097 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2098 #[test]
2099 fn intersect_numeric_stable() {
2100 let a = Tensor::new(vec![4.0, 2.0, 4.0, 1.0, 3.0], vec![5, 1]).unwrap();
2101 let b = Tensor::new(vec![3.0, 4.0, 5.0, 1.0], vec![4, 1]).unwrap();
2102 let eval = intersect_numeric_elements(
2103 a,
2104 b,
2105 &IntersectOptions {
2106 rows: false,
2107 order: IntersectOrder::Stable,
2108 },
2109 )
2110 .expect("intersect");
2111 let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2112 assert_eq!(values.materialize_f64(), vec![4.0, 1.0, 3.0]);
2113 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2114 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2115 assert_eq!(ia.materialize_f64(), vec![1.0, 4.0, 5.0]);
2116 assert_eq!(ib.materialize_f64(), vec![2.0, 4.0, 1.0]);
2117 }
2118
2119 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2120 #[test]
2121 fn intersect_numeric_handles_nan() {
2122 let a = Tensor::new(vec![f64::NAN, 1.0, f64::NAN], vec![3, 1]).unwrap();
2123 let b = Tensor::new(vec![2.0, f64::NAN], vec![2, 1]).unwrap();
2124 let eval = intersect_numeric_elements(
2125 a,
2126 b,
2127 &IntersectOptions {
2128 rows: false,
2129 order: IntersectOrder::Sorted,
2130 },
2131 )
2132 .expect("intersect");
2133 let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2134 assert_eq!(values.materialize_f64().len(), 1);
2135 assert!(values.materialize_f64()[0].is_nan());
2136 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2137 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2138 assert_eq!(ia.materialize_f64(), vec![1.0]);
2139 assert_eq!(ib.materialize_f64(), vec![2.0]);
2140 }
2141
2142 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2143 #[test]
2144 fn intersect_complex_with_real_inputs() {
2145 let complex =
2146 ComplexTensor::new(vec![(1.0, 0.0), (2.0, 0.0), (3.0, 1.0)], vec![3, 1]).unwrap();
2147 let real = Tensor::new(vec![2.0, 4.0, 1.0], vec![3, 1]).unwrap();
2148 let real_complex = tensor_to_complex_owned("intersect", real).unwrap();
2149 let eval = intersect_complex(
2150 complex,
2151 real_complex,
2152 &IntersectOptions {
2153 rows: false,
2154 order: IntersectOrder::Sorted,
2155 },
2156 )
2157 .expect("intersect complex");
2158 match eval.values_value() {
2159 Value::ComplexTensor(t) => {
2160 assert_eq!(t.materialize_f64(), vec![(1.0, 0.0), (2.0, 0.0)]);
2161 }
2162 other => panic!("expected complex tensor, got {other:?}"),
2163 }
2164 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2165 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2166 assert_eq!(ia.materialize_f64(), vec![1.0, 2.0]);
2167 assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2168 }
2169
2170 #[test]
2171 fn intersect_complex_real_alignment_reads_typed_integer_storage_exactly() {
2172 let real =
2173 Tensor::new_integer(IntegerStorage::I64(vec![i64::MIN, -7, 3]), vec![3, 1]).unwrap();
2174 let complex = ComplexTensor::new(vec![(-7.0, 0.0), (4.0, 0.0)], vec![2, 1]).unwrap();
2175
2176 let eval = evaluate_sync(Value::Tensor(real), Value::ComplexTensor(complex), &[])
2177 .expect("intersect");
2178 let Value::Complex(re, im) = eval.values_value() else {
2179 panic!("expected complex scalar");
2180 };
2181 assert_eq!(re, -7.0);
2182 assert_eq!(im, 0.0);
2183 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2184 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2185 assert_eq!(ia.materialize_f64(), vec![2.0]);
2186 assert_eq!(ib.materialize_f64(), vec![1.0]);
2187 }
2188
2189 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2190 #[test]
2191 fn intersect_numeric_rows_default() {
2192 let a = Tensor::new(vec![1.0, 3.0, 1.0, 2.0, 4.0, 2.0], vec![3, 2]).unwrap();
2193 let b = Tensor::new(vec![1.0, 5.0, 2.0, 6.0], vec![2, 2]).unwrap();
2194 let eval = intersect_numeric_rows(
2195 a,
2196 b,
2197 &IntersectOptions {
2198 rows: true,
2199 order: IntersectOrder::Sorted,
2200 },
2201 )
2202 .expect("intersect rows");
2203 let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2204 assert_eq!(values.shape, vec![1, 2]);
2205 assert_eq!(values.materialize_f64(), vec![1.0, 2.0]);
2206 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2207 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2208 assert_eq!(ia.materialize_f64(), vec![1.0]);
2209 assert_eq!(ib.materialize_f64(), vec![1.0]);
2210 }
2211
2212 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2213 #[test]
2214 fn intersect_char_elements_basic() {
2215 let a = CharArray::new("cab".chars().collect(), 1, 3).unwrap();
2216 let b = CharArray::new("bcd".chars().collect(), 1, 3).unwrap();
2217 assert_eq!(find_char_index(&b, 'b'), Some(0));
2218 assert_eq!(find_char_index(&b, 'c'), Some(1));
2219 let b_for_eval = CharArray::new("bcd".chars().collect(), 1, 3).unwrap();
2220 let eval = intersect_char_elements(
2221 a,
2222 b_for_eval,
2223 &IntersectOptions {
2224 rows: false,
2225 order: IntersectOrder::Sorted,
2226 },
2227 )
2228 .expect("intersect char");
2229 match eval.values_value() {
2230 Value::CharArray(arr) => {
2231 assert_eq!(arr.rows, 2);
2232 assert_eq!(arr.cols, 1);
2233 assert_eq!(arr.data, vec!['b', 'c']);
2234 }
2235 other => panic!("expected char array, got {other:?}"),
2236 }
2237 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2238 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2239 assert_eq!(ia.materialize_f64(), vec![3.0, 1.0]);
2240 assert_eq!(ib.materialize_f64(), vec![1.0, 2.0]);
2241 }
2242
2243 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2244 #[test]
2245 fn intersect_string_elements_stable() {
2246 let a = StringArray::new(
2247 vec!["apple".into(), "orange".into(), "pear".into()],
2248 vec![3, 1],
2249 )
2250 .unwrap();
2251 let b = StringArray::new(
2252 vec!["pear".into(), "grape".into(), "orange".into()],
2253 vec![3, 1],
2254 )
2255 .unwrap();
2256 let eval = intersect_string_elements(
2257 a,
2258 b,
2259 &IntersectOptions {
2260 rows: false,
2261 order: IntersectOrder::Stable,
2262 },
2263 )
2264 .expect("intersect string");
2265 match eval.values_value() {
2266 Value::StringArray(arr) => {
2267 assert_eq!(arr.shape, vec![2, 1]);
2268 assert_eq!(arr.data, vec!["orange".to_string(), "pear".to_string()]);
2269 }
2270 other => panic!("expected string array, got {other:?}"),
2271 }
2272 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2273 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2274 assert_eq!(ia.materialize_f64(), vec![2.0, 3.0]);
2275 assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2276 }
2277
2278 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2279 #[test]
2280 fn intersect_rejects_legacy_option() {
2281 let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2282 let err = evaluate_sync(
2283 Value::Tensor(tensor.clone()),
2284 Value::Tensor(tensor),
2285 &[Value::from("legacy")],
2286 )
2287 .unwrap_err();
2288 assert_eq!(
2289 err.identifier(),
2290 INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED.identifier
2291 );
2292 }
2293
2294 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2295 #[test]
2296 fn intersect_rejects_conflicting_order_options() {
2297 let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2298 let err = evaluate_sync(
2299 Value::Tensor(tensor.clone()),
2300 Value::Tensor(tensor),
2301 &[Value::from("stable"), Value::from("sorted")],
2302 )
2303 .unwrap_err();
2304 assert_eq!(
2305 err.identifier(),
2306 INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS.identifier
2307 );
2308 }
2309
2310 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2311 #[test]
2312 fn intersect_rejects_unknown_option() {
2313 let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2314 let err = evaluate_sync(
2315 Value::Tensor(tensor.clone()),
2316 Value::Tensor(tensor),
2317 &[Value::from("bogus")],
2318 )
2319 .unwrap_err();
2320 assert_eq!(err.identifier(), INTERSECT_ERROR_UNKNOWN_OPTION.identifier);
2321 }
2322
2323 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2324 #[test]
2325 fn intersect_rows_dimension_mismatch() {
2326 let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
2327 let b = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2328 let err = intersect_numeric_rows(
2329 a,
2330 b,
2331 &IntersectOptions {
2332 rows: true,
2333 order: IntersectOrder::Sorted,
2334 },
2335 )
2336 .unwrap_err();
2337 assert_eq!(
2338 err.identifier(),
2339 INTERSECT_ERROR_ROWS_COLUMN_MISMATCH.identifier
2340 );
2341 }
2342
2343 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2344 #[test]
2345 fn intersect_mixed_types_error() {
2346 let a = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2347 let b = CharArray::new(vec!['a', 'b'], 1, 2).unwrap();
2348 let err = intersect_host(
2349 Value::Tensor(a),
2350 Value::CharArray(b),
2351 &IntersectOptions {
2352 rows: false,
2353 order: IntersectOrder::Sorted,
2354 },
2355 )
2356 .unwrap_err();
2357 assert_eq!(
2358 err.identifier(),
2359 INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE.identifier
2360 );
2361 }
2362
2363 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2364 #[test]
2365 fn intersect_gpu_roundtrip() {
2366 test_support::with_test_provider(|provider| {
2367 let a = Tensor::new(vec![4.0, 1.0, 2.0, 1.0], vec![4, 1]).unwrap();
2368 let b = Tensor::new(vec![2.0, 5.0, 1.0], vec![3, 1]).unwrap();
2369 let view_a = HostTensorView {
2370 data: &a.materialize_f64(),
2371 shape: &a.shape,
2372 };
2373 let view_b = HostTensorView {
2374 data: &b.materialize_f64(),
2375 shape: &b.shape,
2376 };
2377 let handle_a = provider.upload(&view_a).expect("upload A");
2378 let handle_b = provider.upload(&view_b).expect("upload B");
2379 let eval = evaluate_sync(Value::GpuTensor(handle_a), Value::GpuTensor(handle_b), &[])
2380 .expect("intersect");
2381 let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2382 assert_eq!(values.materialize_f64(), vec![1.0, 2.0]);
2383 let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2384 let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2385 assert_eq!(ia.materialize_f64(), vec![2.0, 3.0]);
2386 assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2387 });
2388 }
2389
2390 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2391 #[test]
2392 fn intersect_two_outputs_from_evaluate() {
2393 let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2394 let b = Tensor::new(vec![3.0, 1.0], vec![2, 1]).unwrap();
2395 let eval = intersect_numeric_elements(
2396 a,
2397 b,
2398 &IntersectOptions {
2399 rows: false,
2400 order: IntersectOrder::Sorted,
2401 },
2402 )
2403 .unwrap();
2404 let (_c, ia) = eval.clone().into_pair();
2405 let ia_tensor = tensor::value_into_tensor_for("intersect", ia).unwrap();
2406 assert_eq!(ia_tensor.materialize_f64(), vec![1.0, 3.0]);
2407 let (_c, ia2, ib2) = eval.into_triple();
2408 let ia_tensor2 = tensor::value_into_tensor_for("intersect", ia2).unwrap();
2409 let ib_tensor2 = tensor::value_into_tensor_for("intersect", ib2).unwrap();
2410 assert_eq!(ia_tensor2.materialize_f64(), vec![1.0, 3.0]);
2411 assert_eq!(ib_tensor2.materialize_f64(), vec![2.0, 1.0]);
2412 }
2413
2414 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2415 #[test]
2416 #[cfg(feature = "wgpu")]
2417 fn intersect_wgpu_matches_cpu() {
2418 let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2419 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2420 );
2421 let a = Tensor::new(vec![4.0, 1.0, 2.0, 3.0], vec![4, 1]).unwrap();
2422 let b = Tensor::new(vec![2.0, 6.0, 3.0], vec![3, 1]).unwrap();
2423
2424 let cpu_eval = intersect_numeric_elements(
2425 a.clone(),
2426 b.clone(),
2427 &IntersectOptions {
2428 rows: false,
2429 order: IntersectOrder::Sorted,
2430 },
2431 )
2432 .unwrap();
2433 let cpu_values =
2434 tensor::value_into_tensor_for("intersect", cpu_eval.values_value()).unwrap();
2435 let cpu_ia = tensor::value_into_tensor_for("intersect", cpu_eval.ia_value()).unwrap();
2436 let cpu_ib = tensor::value_into_tensor_for("intersect", cpu_eval.ib_value()).unwrap();
2437
2438 let provider = runmat_accelerate_api::provider().expect("provider");
2439 let view_a = HostTensorView {
2440 data: &a.materialize_f64(),
2441 shape: &a.shape,
2442 };
2443 let view_b = HostTensorView {
2444 data: &b.materialize_f64(),
2445 shape: &b.shape,
2446 };
2447 let handle_a = provider.upload(&view_a).expect("upload A");
2448 let handle_b = provider.upload(&view_b).expect("upload B");
2449 let gpu_eval = evaluate_sync(Value::GpuTensor(handle_a), Value::GpuTensor(handle_b), &[])
2450 .expect("intersect");
2451 let gpu_values =
2452 tensor::value_into_tensor_for("intersect", gpu_eval.values_value()).unwrap();
2453 let gpu_ia = tensor::value_into_tensor_for("intersect", gpu_eval.ia_value()).unwrap();
2454 let gpu_ib = tensor::value_into_tensor_for("intersect", gpu_eval.ib_value()).unwrap();
2455
2456 assert_eq!(gpu_values.materialize_f64(), cpu_values.materialize_f64());
2457 assert_eq!(gpu_ia.materialize_f64(), cpu_ia.materialize_f64());
2458 assert_eq!(gpu_ib.materialize_f64(), cpu_ib.materialize_f64());
2459 }
2460}