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