Skip to main content

runmat_runtime/builtins/array/sorting_sets/
setdiff.rs

1//! MATLAB-compatible `setdiff` builtin with GPU-aware semantics for RunMat.
2//!
3//! Provides element-wise and row-wise set difference with optional stable
4//! ordering. GPU tensors use a provider hook or typed host fallback, then public
5//! outputs are restored to the owning provider.
6
7use 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
313/// Evaluate the `setdiff` builtin once and expose all outputs.
314pub 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                // Fall back to host gather when provider does not support setdiff.
421            }
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
646/// Helper exposed for acceleration providers handling numeric tensors entirely on the host.
647pub 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}