1use std::cmp::Ordering;
8use std::collections::{hash_map::Entry, HashMap};
9
10use runmat_accelerate_api::GpuTensorHandle;
11use runmat_builtins::{
12 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
13 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
14 CharArray, ComplexTensor, NumericDType, StringArray, Tensor, Value,
15};
16use runmat_macros::runtime_builtin;
17
18use super::type_resolvers::set_values_output_type;
19use crate::build_runtime_error;
20use crate::builtins::common::arg_tokens::tokens_from_values;
21use crate::builtins::common::gpu_helpers;
22use crate::builtins::common::random_args::complex_tensor_into_value;
23use crate::builtins::common::spec::{
24 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
25 ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
26};
27use crate::builtins::common::tensor;
28
29#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::array::sorting_sets::setxor")]
30pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
31 name: "setxor",
32 op_kind: GpuOpKind::Custom("setxor"),
33 supported_precisions: &[ScalarType::F32, ScalarType::F64],
34 broadcast: BroadcastSemantics::None,
35 provider_hooks: &[],
36 constant_strategy: ConstantStrategy::InlineLiteral,
37 residency: ResidencyPolicy::GatherImmediately,
38 nan_mode: ReductionNaN::Include,
39 two_pass_threshold: None,
40 workgroup_size: None,
41 accepts_nan_mode: true,
42 notes: "`setxor` currently gathers GPU tensors and evaluates on the host.",
43};
44
45#[runmat_macros::register_fusion_spec(
46 builtin_path = "crate::builtins::array::sorting_sets::setxor"
47)]
48pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
49 name: "setxor",
50 shape: ShapeRequirements::Any,
51 constant_strategy: ConstantStrategy::InlineLiteral,
52 elementwise: None,
53 reduction: None,
54 emits_nan: true,
55 notes: "`setxor` terminates fusion chains and materialises results on the host; upstream tensors are gathered when necessary.",
56};
57
58const BUILTIN_NAME: &str = "setxor";
59
60const SETXOR_OUTPUT_C: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
61 name: "C",
62 ty: BuiltinParamType::Any,
63 arity: BuiltinParamArity::Required,
64 default: None,
65 description: "Values or rows that appear in exactly one input.",
66}];
67
68const SETXOR_OUTPUT_C_IA_IB: [BuiltinParamDescriptor; 3] = [
69 BuiltinParamDescriptor {
70 name: "C",
71 ty: BuiltinParamType::Any,
72 arity: BuiltinParamArity::Required,
73 default: None,
74 description: "Values or rows that appear in exactly one input.",
75 },
76 BuiltinParamDescriptor {
77 name: "ia",
78 ty: BuiltinParamType::NumericArray,
79 arity: BuiltinParamArity::Required,
80 default: None,
81 description: "Indices selecting values or rows from A.",
82 },
83 BuiltinParamDescriptor {
84 name: "ib",
85 ty: BuiltinParamType::NumericArray,
86 arity: BuiltinParamArity::Required,
87 default: None,
88 description: "Indices selecting values or rows from B.",
89 },
90];
91
92const SETXOR_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 SETXOR_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 SETXOR_SIGNATURES: [BuiltinSignatureDescriptor; 4] = [
134 BuiltinSignatureDescriptor {
135 label: "C = setxor(A, B)",
136 inputs: &SETXOR_INPUTS_A_B,
137 outputs: &SETXOR_OUTPUT_C,
138 },
139 BuiltinSignatureDescriptor {
140 label: "C = setxor(A, B, option...)",
141 inputs: &SETXOR_INPUTS_A_B_OPTIONS,
142 outputs: &SETXOR_OUTPUT_C,
143 },
144 BuiltinSignatureDescriptor {
145 label: "[C, ia, ib] = setxor(A, B)",
146 inputs: &SETXOR_INPUTS_A_B,
147 outputs: &SETXOR_OUTPUT_C_IA_IB,
148 },
149 BuiltinSignatureDescriptor {
150 label: "[C, ia, ib] = setxor(A, B, option...)",
151 inputs: &SETXOR_INPUTS_A_B_OPTIONS,
152 outputs: &SETXOR_OUTPUT_C_IA_IB,
153 },
154];
155
156const SETXOR_ERROR_LEGACY_OPTION_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
157 code: "RM.SETXOR.LEGACY_OPTION_UNSUPPORTED",
158 identifier: Some("RunMat:setxor:LegacyOptionUnsupported"),
159 when: "Legacy compatibility options are requested.",
160 message: "setxor: the 'legacy' behaviour is not supported",
161};
162
163const SETXOR_ERROR_CONFLICTING_ORDER_OPTIONS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
164 code: "RM.SETXOR.CONFLICTING_ORDER_OPTIONS",
165 identifier: Some("RunMat:setxor:ConflictingOrderOptions"),
166 when: "Both 'sorted' and 'stable' options are provided.",
167 message: "setxor: cannot combine 'sorted' with 'stable'",
168};
169
170const SETXOR_ERROR_UNKNOWN_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
171 code: "RM.SETXOR.UNKNOWN_OPTION",
172 identifier: Some("RunMat:setxor:UnknownOption"),
173 when: "An unsupported option token is provided.",
174 message: "setxor: unrecognised option",
175};
176
177const SETXOR_ERROR_ROWS_COLUMN_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
178 code: "RM.SETXOR.ROWS_COLUMN_MISMATCH",
179 identifier: Some("RunMat:setxor:RowsColumnMismatch"),
180 when: "'rows' mode is used and column counts differ.",
181 message: "setxor: inputs must have the same number of columns when using 'rows'",
182};
183
184const SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
185 code: "RM.SETXOR.UNSUPPORTED_INPUT_TYPE",
186 identifier: Some("RunMat:setxor:UnsupportedInputType"),
187 when: "Input values cannot be converted into supported setxor domains.",
188 message: "setxor: unsupported input type",
189};
190
191const SETXOR_ERROR_NUMERIC_CLASS_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
192 code: "RM.SETXOR.NUMERIC_CLASS_MISMATCH",
193 identifier: Some("RunMat:setxor:NumericClassMismatch"),
194 when: "Numeric inputs have incompatible nondouble classes.",
195 message: "setxor: numeric inputs must have the same class, except double may be combined with one nondouble class",
196};
197
198const SETXOR_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
199 code: "RM.SETXOR.INVALID_ARGUMENT",
200 identifier: Some("RunMat:setxor:InvalidArgument"),
201 when: "Option arguments are not string-like where required.",
202 message: "setxor: expected string option arguments",
203};
204
205const SETXOR_ERROR_TOO_MANY_OUTPUTS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
206 code: "RM.SETXOR.TOO_MANY_OUTPUTS",
207 identifier: Some("RunMat:setxor:TooManyOutputs"),
208 when: "More than three output arguments are requested.",
209 message: "setxor: too many output arguments",
210};
211
212const SETXOR_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
213 code: "RM.SETXOR.INTERNAL",
214 identifier: Some("RunMat:setxor:Internal"),
215 when: "Internal conversion, allocation, or provider decode fails.",
216 message: "setxor: internal operation failed",
217};
218
219const SETXOR_ERRORS: [BuiltinErrorDescriptor; 9] = [
220 SETXOR_ERROR_LEGACY_OPTION_UNSUPPORTED,
221 SETXOR_ERROR_CONFLICTING_ORDER_OPTIONS,
222 SETXOR_ERROR_UNKNOWN_OPTION,
223 SETXOR_ERROR_ROWS_COLUMN_MISMATCH,
224 SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE,
225 SETXOR_ERROR_NUMERIC_CLASS_MISMATCH,
226 SETXOR_ERROR_INVALID_ARGUMENT,
227 SETXOR_ERROR_TOO_MANY_OUTPUTS,
228 SETXOR_ERROR_INTERNAL,
229];
230
231pub const SETXOR_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
232 signatures: &SETXOR_SIGNATURES,
233 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
234 completion_policy: BuiltinCompletionPolicy::Public,
235 errors: &SETXOR_ERRORS,
236};
237
238#[derive(Debug, Clone, Copy, PartialEq, Eq)]
239enum SetxorOrder {
240 Sorted,
241 Stable,
242}
243
244#[derive(Debug, Clone)]
245struct SetxorOptions {
246 rows: bool,
247 order: SetxorOrder,
248}
249
250#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
251enum Origin {
252 A,
253 B,
254}
255
256#[derive(Debug)]
257struct SymEntry<T> {
258 value: T,
259 a_index: Option<usize>,
260 b_index: Option<usize>,
261 order_rank: usize,
262}
263
264#[derive(Clone, Debug)]
265struct ElementMeta {
266 row_output: bool,
267 dtype: NumericDType,
268}
269
270#[derive(Debug, Clone, PartialEq, Eq, Hash)]
271enum NumericKey {
272 Value(u64),
273 UniqueNan(Origin, usize),
274}
275
276#[derive(Debug, Clone, PartialEq, Eq, Hash)]
277enum NumericRowKey {
278 Values(Vec<u64>),
279 UniqueNan(Origin, usize),
280}
281
282#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
283struct ComplexKey {
284 re: u64,
285 im: u64,
286}
287
288#[derive(Debug, Clone, PartialEq, Eq, Hash)]
289enum ComplexElementKey {
290 Value(ComplexKey),
291 UniqueNan(Origin, usize),
292}
293
294#[derive(Debug, Clone, PartialEq, Eq, Hash)]
295enum ComplexRowKey {
296 Values(Vec<ComplexKey>),
297 UniqueNan(Origin, usize),
298}
299
300#[derive(Debug, Clone, PartialEq, Eq, Hash)]
301struct RowCharKey(Vec<u32>);
302
303#[derive(Debug, Clone, PartialEq, Eq, Hash)]
304struct RowStringKey(Vec<String>);
305
306fn setxor_error_with(
307 error: &'static BuiltinErrorDescriptor,
308 message: impl Into<String>,
309) -> crate::RuntimeError {
310 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
311 if let Some(identifier) = error.identifier {
312 builder = builder.with_identifier(identifier);
313 }
314 builder.build()
315}
316
317fn setxor_error(error: &'static BuiltinErrorDescriptor) -> crate::RuntimeError {
318 setxor_error_with(error, error.message)
319}
320
321fn setxor_internal_error(message: impl Into<String>) -> crate::RuntimeError {
322 setxor_error_with(&SETXOR_ERROR_INTERNAL, message)
323}
324
325#[runtime_builtin(
326 name = "setxor",
327 category = "array/sorting_sets",
328 summary = "Return the symmetric difference of two arrays or row sets.",
329 keywords = "setxor,symmetric difference,exclusive or,stable,rows,indices,gpu",
330 accel = "array_construct",
331 sink = true,
332 type_resolver(set_values_output_type),
333 descriptor(crate::builtins::array::sorting_sets::setxor::SETXOR_DESCRIPTOR),
334 builtin_path = "crate::builtins::array::sorting_sets::setxor"
335)]
336async fn setxor_builtin(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
337 if matches!(crate::output_count::current_output_count(), Some(n) if n > 3) {
338 return Err(setxor_error_with(
339 &SETXOR_ERROR_TOO_MANY_OUTPUTS,
340 "setxor: too many output arguments; maximum is 3",
341 ));
342 }
343 let eval = evaluate(a, b, &rest).await?;
344 if let Some(out_count) = crate::output_count::current_output_count() {
345 if out_count == 0 {
346 return Ok(Value::OutputList(Vec::new()));
347 }
348 if out_count == 1 {
349 return Ok(Value::OutputList(vec![eval.into_values_value()]));
350 }
351 let (values, ia, ib) = eval.into_triple();
352 return Ok(crate::output_count::output_list_with_padding(
353 out_count,
354 vec![values, ia, ib],
355 ));
356 }
357 Ok(eval.into_values_value())
358}
359
360pub async fn evaluate(
361 a: Value,
362 b: Value,
363 rest: &[Value],
364) -> crate::BuiltinResult<SetxorEvaluation> {
365 let opts = parse_options(rest)?;
366 match (a, b) {
367 (Value::GpuTensor(handle_a), Value::GpuTensor(handle_b)) => {
368 setxor_gpu_pair(handle_a, handle_b, &opts).await
369 }
370 (Value::GpuTensor(handle_a), other) => setxor_gpu_mixed(handle_a, other, &opts, true).await,
371 (other, Value::GpuTensor(handle_b)) => {
372 setxor_gpu_mixed(handle_b, other, &opts, false).await
373 }
374 (left, right) => setxor_host(left, right, &opts),
375 }
376}
377
378fn parse_options(rest: &[Value]) -> crate::BuiltinResult<SetxorOptions> {
379 let mut opts = SetxorOptions {
380 rows: false,
381 order: SetxorOrder::Sorted,
382 };
383 let mut seen_order: Option<SetxorOrder> = None;
384
385 let tokens = tokens_from_values(rest);
386 for (arg, token) in rest.iter().zip(tokens.iter()) {
387 let text = match token {
388 crate::builtins::common::arg_tokens::ArgToken::String(text) => text.as_str(),
389 _ => {
390 let text = tensor::value_to_string(arg)
391 .ok_or_else(|| setxor_error(&SETXOR_ERROR_INVALID_ARGUMENT))?;
392 let lowered = text.trim().to_ascii_lowercase();
393 parse_setxor_option(&mut opts, &mut seen_order, &lowered)?;
394 continue;
395 }
396 };
397 parse_setxor_option(&mut opts, &mut seen_order, text)?;
398 }
399
400 Ok(opts)
401}
402
403fn parse_setxor_option(
404 opts: &mut SetxorOptions,
405 seen_order: &mut Option<SetxorOrder>,
406 lowered: &str,
407) -> crate::BuiltinResult<()> {
408 match lowered {
409 "rows" => opts.rows = true,
410 "sorted" => {
411 if let Some(prev) = seen_order {
412 if *prev != SetxorOrder::Sorted {
413 return Err(setxor_error(&SETXOR_ERROR_CONFLICTING_ORDER_OPTIONS));
414 }
415 }
416 *seen_order = Some(SetxorOrder::Sorted);
417 opts.order = SetxorOrder::Sorted;
418 }
419 "stable" => {
420 if let Some(prev) = seen_order {
421 if *prev != SetxorOrder::Stable {
422 return Err(setxor_error(&SETXOR_ERROR_CONFLICTING_ORDER_OPTIONS));
423 }
424 }
425 *seen_order = Some(SetxorOrder::Stable);
426 opts.order = SetxorOrder::Stable;
427 }
428 "legacy" | "r2012a" => {
429 return Err(setxor_error(&SETXOR_ERROR_LEGACY_OPTION_UNSUPPORTED));
430 }
431 other => {
432 return Err(setxor_error_with(
433 &SETXOR_ERROR_UNKNOWN_OPTION,
434 format!("setxor: unrecognised option '{other}'"),
435 ))
436 }
437 }
438 Ok(())
439}
440
441async fn setxor_gpu_pair(
442 handle_a: GpuTensorHandle,
443 handle_b: GpuTensorHandle,
444 opts: &SetxorOptions,
445) -> crate::BuiltinResult<SetxorEvaluation> {
446 let tensor_a = gpu_helpers::gather_tensor_async(&handle_a).await?;
447 let tensor_b = gpu_helpers::gather_tensor_async(&handle_b).await?;
448 setxor_numeric(tensor_a, tensor_b, opts)
449}
450
451async fn setxor_gpu_mixed(
452 handle_gpu: GpuTensorHandle,
453 other: Value,
454 opts: &SetxorOptions,
455 gpu_is_a: bool,
456) -> crate::BuiltinResult<SetxorEvaluation> {
457 let tensor_gpu = gpu_helpers::gather_tensor_async(&handle_gpu).await?;
458 if matches!(other, Value::ComplexTensor(_) | Value::Complex(_, _)) {
459 let complex_gpu = tensor_to_complex(tensor_gpu)?;
460 let complex_other = value_into_complex_tensor(other)?;
461 return if gpu_is_a {
462 setxor_complex(complex_gpu, complex_other, opts)
463 } else {
464 setxor_complex(complex_other, complex_gpu, opts)
465 };
466 }
467 let tensor_other =
468 tensor::value_into_tensor_for("setxor", other).map_err(setxor_internal_error)?;
469 if gpu_is_a {
470 setxor_numeric(tensor_gpu, tensor_other, opts)
471 } else {
472 setxor_numeric(tensor_other, tensor_gpu, opts)
473 }
474}
475
476fn setxor_host(a: Value, b: Value, opts: &SetxorOptions) -> crate::BuiltinResult<SetxorEvaluation> {
477 match (a, b) {
478 (Value::ComplexTensor(at), right) => {
479 let bt = value_into_complex_tensor(right)?;
480 setxor_complex(at, bt, opts)
481 }
482 (left, Value::ComplexTensor(bt)) => {
483 let at = value_into_complex_tensor(left)?;
484 setxor_complex(at, bt, opts)
485 }
486 (Value::Complex(re, im), right) => {
487 let at = ComplexTensor::new(vec![(re, im)], vec![1, 1])
488 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
489 let bt = value_into_complex_tensor(right)?;
490 setxor_complex(at, bt, opts)
491 }
492 (left, Value::Complex(re, im)) => {
493 let at = value_into_complex_tensor(left)?;
494 let bt = ComplexTensor::new(vec![(re, im)], vec![1, 1])
495 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
496 setxor_complex(at, bt, opts)
497 }
498 (Value::CharArray(ac), Value::CharArray(bc)) => setxor_char(ac, bc, opts),
499 (Value::StringArray(astring), right) if value_is_string_compatible(&right) => {
500 let bstring = value_into_string_array(right)?;
501 setxor_string(astring, bstring, opts)
502 }
503 (left, Value::StringArray(bstring)) if value_is_string_compatible(&left) => {
504 let astring = value_into_string_array(left)?;
505 setxor_string(astring, bstring, opts)
506 }
507 (Value::String(a), right) if value_is_string_compatible(&right) => {
508 let astring = StringArray::new(vec![a], vec![1, 1])
509 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
510 let bstring = value_into_string_array(right)?;
511 setxor_string(astring, bstring, opts)
512 }
513 (left, Value::String(b)) if value_is_string_compatible(&left) => {
514 let astring = value_into_string_array(left)?;
515 let bstring = StringArray::new(vec![b], vec![1, 1])
516 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
517 setxor_string(astring, bstring, opts)
518 }
519 (Value::StringArray(astring), Value::StringArray(bstring)) => {
520 setxor_string(astring, bstring, opts)
521 }
522 (Value::StringArray(astring), Value::String(b)) => {
523 let bstring = StringArray::new(vec![b], vec![1, 1])
524 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
525 setxor_string(astring, bstring, opts)
526 }
527 (Value::String(a), Value::StringArray(bstring)) => {
528 let astring = StringArray::new(vec![a], vec![1, 1])
529 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
530 setxor_string(astring, bstring, opts)
531 }
532 (Value::String(a), Value::String(b)) => {
533 let astring = StringArray::new(vec![a], vec![1, 1])
534 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
535 let bstring = StringArray::new(vec![b], vec![1, 1])
536 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
537 setxor_string(astring, bstring, opts)
538 }
539 (Value::CharArray(ac), right) if value_is_char_numeric_compatible(&right) => {
540 let bc = value_into_char_array(right)?;
541 setxor_char(ac, bc, opts)
542 }
543 (left, Value::CharArray(bc)) if value_is_char_numeric_compatible(&left) => {
544 let ac = value_into_char_array(left)?;
545 setxor_char(ac, bc, opts)
546 }
547 (left, right) => {
548 let tensor_a = tensor::value_into_tensor_for("setxor", left)
549 .map_err(|e| setxor_error_with(&SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
550 let tensor_b = tensor::value_into_tensor_for("setxor", right)
551 .map_err(|e| setxor_error_with(&SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
552 setxor_numeric(tensor_a, tensor_b, opts)
553 }
554 }
555}
556
557fn value_into_complex_tensor(value: Value) -> crate::BuiltinResult<ComplexTensor> {
558 match value {
559 Value::ComplexTensor(tensor) => Ok(tensor),
560 Value::Complex(re, im) => ComplexTensor::new(vec![(re, im)], vec![1, 1])
561 .map_err(|e| setxor_internal_error(format!("setxor: {e}"))),
562 other => {
563 let tensor = tensor::value_into_tensor_for("setxor", other)
564 .map_err(|e| setxor_error_with(&SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
565 tensor_to_complex(tensor)
566 }
567 }
568}
569
570fn tensor_to_complex(tensor: Tensor) -> crate::BuiltinResult<ComplexTensor> {
571 let shape = tensor.shape;
572 let data = tensor
573 .data
574 .into_iter()
575 .map(|real| (real, 0.0))
576 .collect::<Vec<_>>();
577 ComplexTensor::new(data, shape).map_err(|e| setxor_internal_error(format!("setxor: {e}")))
578}
579
580fn value_is_string_compatible(value: &Value) -> bool {
581 matches!(
582 value,
583 Value::StringArray(_) | Value::String(_) | Value::CharArray(_)
584 )
585}
586
587fn value_into_string_array(value: Value) -> crate::BuiltinResult<StringArray> {
588 match value {
589 Value::StringArray(array) => Ok(array),
590 Value::String(value) => StringArray::new(vec![value], vec![1, 1])
591 .map_err(|e| setxor_internal_error(format!("setxor: {e}"))),
592 Value::CharArray(chars) => char_array_to_string_array(chars),
593 other => Err(setxor_error_with(
594 &SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE,
595 format!("setxor: cannot convert {other:?} to string array"),
596 )),
597 }
598}
599
600fn char_array_to_string_array(chars: CharArray) -> crate::BuiltinResult<StringArray> {
601 let values = (0..chars.rows)
602 .map(|row| {
603 chars.data[row * chars.cols..row * chars.cols + chars.cols]
604 .iter()
605 .collect()
606 })
607 .collect::<Vec<String>>();
608 let shape = if chars.rows == 0 {
609 vec![0, 1]
610 } else if chars.rows == 1 {
611 vec![1, 1]
612 } else {
613 vec![chars.rows, 1]
614 };
615 StringArray::new(values, shape).map_err(|e| setxor_internal_error(format!("setxor: {e}")))
616}
617
618fn value_is_char_numeric_compatible(value: &Value) -> bool {
619 matches!(
620 value,
621 Value::Tensor(_) | Value::LogicalArray(_) | Value::Num(_) | Value::Int(_) | Value::Bool(_)
622 )
623}
624
625fn value_into_char_array(value: Value) -> crate::BuiltinResult<CharArray> {
626 let tensor = tensor::value_into_tensor_for("setxor", value)
627 .map_err(|e| setxor_error_with(&SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
628 tensor_into_char_array(tensor)
629}
630
631fn tensor_into_char_array(tensor: Tensor) -> crate::BuiltinResult<CharArray> {
632 let rows = tensor.rows;
633 let cols = tensor.cols;
634 let mut values = vec!['\0'; rows * cols];
635 for col in 0..cols {
636 for row in 0..rows {
637 let value = tensor.data[row + col * rows];
638 values[row * cols + col] = f64_to_char(value)?;
639 }
640 }
641 CharArray::new(values, rows, cols).map_err(|e| setxor_internal_error(format!("setxor: {e}")))
642}
643
644fn f64_to_char(value: f64) -> crate::BuiltinResult<char> {
645 if !value.is_finite() || value.fract() != 0.0 || value < 0.0 || value > u32::MAX as f64 {
646 return Err(setxor_error_with(
647 &SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE,
648 "setxor: numeric values mixed with char inputs must be finite character codes",
649 ));
650 }
651 char::from_u32(value as u32).ok_or_else(|| {
652 setxor_error_with(
653 &SETXOR_ERROR_UNSUPPORTED_INPUT_TYPE,
654 "setxor: numeric values mixed with char inputs must be valid character codes",
655 )
656 })
657}
658
659fn setxor_numeric(
660 a: Tensor,
661 b: Tensor,
662 opts: &SetxorOptions,
663) -> crate::BuiltinResult<SetxorEvaluation> {
664 if opts.rows {
665 setxor_numeric_rows(a, b, opts)
666 } else {
667 let meta = element_meta(&a.shape, a.dtype, &b.shape, b.dtype)?;
668 let mut entries = Vec::<SymEntry<f64>>::new();
669 let mut map: HashMap<NumericKey, usize> = HashMap::new();
670 let mut order_counter = 0usize;
671 for (idx, &value) in a.data.iter().enumerate() {
672 add_sym_entry(
673 &mut entries,
674 &mut map,
675 numeric_key(value, Origin::A, idx),
676 value,
677 Origin::A,
678 idx,
679 &mut order_counter,
680 );
681 }
682 for (idx, &value) in b.data.iter().enumerate() {
683 add_sym_entry(
684 &mut entries,
685 &mut map,
686 numeric_key(value, Origin::B, idx),
687 value,
688 Origin::B,
689 idx,
690 &mut order_counter,
691 );
692 }
693 assemble_numeric(entries, opts, &meta)
694 }
695}
696
697fn setxor_numeric_rows(
698 a: Tensor,
699 b: Tensor,
700 opts: &SetxorOptions,
701) -> crate::BuiltinResult<SetxorEvaluation> {
702 if a.shape.len() != 2 || b.shape.len() != 2 {
703 return Err(setxor_internal_error(
704 "setxor: 'rows' option requires 2-D numeric matrices",
705 ));
706 }
707 if a.shape[1] != b.shape[1] {
708 return Err(setxor_error(&SETXOR_ERROR_ROWS_COLUMN_MISMATCH));
709 }
710 let rows_a = a.shape[0];
711 let rows_b = b.shape[0];
712 let cols = a.shape[1];
713 let dtype = numeric_output_dtype(a.dtype, b.dtype)?;
714 let mut entries = Vec::<SymEntry<Vec<f64>>>::new();
715 let mut map: HashMap<NumericRowKey, usize> = HashMap::new();
716 let mut order_counter = 0usize;
717 for row in 0..rows_a {
718 let values = numeric_row(&a, row, cols);
719 let key = numeric_row_key(&values, Origin::A, row);
720 add_sym_entry(
721 &mut entries,
722 &mut map,
723 key,
724 values,
725 Origin::A,
726 row,
727 &mut order_counter,
728 );
729 }
730 for row in 0..rows_b {
731 let values = numeric_row(&b, row, cols);
732 let key = numeric_row_key(&values, Origin::B, row);
733 add_sym_entry(
734 &mut entries,
735 &mut map,
736 key,
737 values,
738 Origin::B,
739 row,
740 &mut order_counter,
741 );
742 }
743 assemble_numeric_rows(entries, opts, cols, dtype)
744}
745
746fn setxor_complex(
747 a: ComplexTensor,
748 b: ComplexTensor,
749 opts: &SetxorOptions,
750) -> crate::BuiltinResult<SetxorEvaluation> {
751 if opts.rows {
752 setxor_complex_rows(a, b, opts)
753 } else {
754 let row_output = element_row_output(&a.shape, &b.shape);
755 let mut entries = Vec::<SymEntry<(f64, f64)>>::new();
756 let mut map: HashMap<ComplexElementKey, usize> = HashMap::new();
757 let mut order_counter = 0usize;
758 for (idx, &value) in a.data.iter().enumerate() {
759 add_sym_entry(
760 &mut entries,
761 &mut map,
762 complex_element_key(value, Origin::A, idx),
763 value,
764 Origin::A,
765 idx,
766 &mut order_counter,
767 );
768 }
769 for (idx, &value) in b.data.iter().enumerate() {
770 add_sym_entry(
771 &mut entries,
772 &mut map,
773 complex_element_key(value, Origin::B, idx),
774 value,
775 Origin::B,
776 idx,
777 &mut order_counter,
778 );
779 }
780 assemble_complex(entries, opts, row_output)
781 }
782}
783
784fn setxor_complex_rows(
785 a: ComplexTensor,
786 b: ComplexTensor,
787 opts: &SetxorOptions,
788) -> crate::BuiltinResult<SetxorEvaluation> {
789 if a.shape.len() != 2 || b.shape.len() != 2 {
790 return Err(setxor_internal_error(
791 "setxor: 'rows' option requires 2-D complex matrices",
792 ));
793 }
794 if a.shape[1] != b.shape[1] {
795 return Err(setxor_error(&SETXOR_ERROR_ROWS_COLUMN_MISMATCH));
796 }
797 let rows_a = a.shape[0];
798 let rows_b = b.shape[0];
799 let cols = a.shape[1];
800 let mut entries = Vec::<SymEntry<Vec<(f64, f64)>>>::new();
801 let mut map: HashMap<ComplexRowKey, usize> = HashMap::new();
802 let mut order_counter = 0usize;
803 for row in 0..rows_a {
804 let values = complex_row(&a, row, cols);
805 let key = complex_row_key(&values, Origin::A, row);
806 add_sym_entry(
807 &mut entries,
808 &mut map,
809 key,
810 values,
811 Origin::A,
812 row,
813 &mut order_counter,
814 );
815 }
816 for row in 0..rows_b {
817 let values = complex_row(&b, row, cols);
818 let key = complex_row_key(&values, Origin::B, row);
819 add_sym_entry(
820 &mut entries,
821 &mut map,
822 key,
823 values,
824 Origin::B,
825 row,
826 &mut order_counter,
827 );
828 }
829 assemble_complex_rows(entries, opts, cols)
830}
831
832fn setxor_char(
833 a: CharArray,
834 b: CharArray,
835 opts: &SetxorOptions,
836) -> crate::BuiltinResult<SetxorEvaluation> {
837 if opts.rows {
838 setxor_char_rows(a, b, opts)
839 } else {
840 let row_output = a.rows == 1 && b.rows == 1;
841 let mut entries = Vec::<SymEntry<char>>::new();
842 let mut map: HashMap<u32, usize> = HashMap::new();
843 let mut order_counter = 0usize;
844 for col in 0..a.cols {
845 for row in 0..a.rows {
846 let linear_idx = row + col * a.rows;
847 let data_idx = row * a.cols + col;
848 let ch = a.data[data_idx];
849 add_sym_entry(
850 &mut entries,
851 &mut map,
852 ch as u32,
853 ch,
854 Origin::A,
855 linear_idx,
856 &mut order_counter,
857 );
858 }
859 }
860 for col in 0..b.cols {
861 for row in 0..b.rows {
862 let linear_idx = row + col * b.rows;
863 let data_idx = row * b.cols + col;
864 let ch = b.data[data_idx];
865 add_sym_entry(
866 &mut entries,
867 &mut map,
868 ch as u32,
869 ch,
870 Origin::B,
871 linear_idx,
872 &mut order_counter,
873 );
874 }
875 }
876 assemble_char(entries, opts, row_output)
877 }
878}
879
880fn setxor_char_rows(
881 a: CharArray,
882 b: CharArray,
883 opts: &SetxorOptions,
884) -> crate::BuiltinResult<SetxorEvaluation> {
885 if a.cols != b.cols {
886 return Err(setxor_error(&SETXOR_ERROR_ROWS_COLUMN_MISMATCH));
887 }
888 let mut entries = Vec::<SymEntry<Vec<char>>>::new();
889 let mut map: HashMap<RowCharKey, usize> = HashMap::new();
890 let mut order_counter = 0usize;
891 for row in 0..a.rows {
892 let values = char_row(&a, row);
893 add_sym_entry(
894 &mut entries,
895 &mut map,
896 RowCharKey(values.iter().map(|&ch| ch as u32).collect()),
897 values,
898 Origin::A,
899 row,
900 &mut order_counter,
901 );
902 }
903 for row in 0..b.rows {
904 let values = char_row(&b, row);
905 add_sym_entry(
906 &mut entries,
907 &mut map,
908 RowCharKey(values.iter().map(|&ch| ch as u32).collect()),
909 values,
910 Origin::B,
911 row,
912 &mut order_counter,
913 );
914 }
915 assemble_char_rows(entries, opts, a.cols)
916}
917
918fn setxor_string(
919 a: StringArray,
920 b: StringArray,
921 opts: &SetxorOptions,
922) -> crate::BuiltinResult<SetxorEvaluation> {
923 if opts.rows {
924 setxor_string_rows(a, b, opts)
925 } else {
926 let row_output = element_row_output(&a.shape, &b.shape);
927 let mut entries = Vec::<SymEntry<String>>::new();
928 let mut map: HashMap<String, usize> = HashMap::new();
929 let mut order_counter = 0usize;
930 for (idx, value) in a.data.iter().enumerate() {
931 add_sym_entry(
932 &mut entries,
933 &mut map,
934 value.clone(),
935 value.clone(),
936 Origin::A,
937 idx,
938 &mut order_counter,
939 );
940 }
941 for (idx, value) in b.data.iter().enumerate() {
942 add_sym_entry(
943 &mut entries,
944 &mut map,
945 value.clone(),
946 value.clone(),
947 Origin::B,
948 idx,
949 &mut order_counter,
950 );
951 }
952 assemble_string(entries, opts, row_output)
953 }
954}
955
956fn setxor_string_rows(
957 a: StringArray,
958 b: StringArray,
959 opts: &SetxorOptions,
960) -> crate::BuiltinResult<SetxorEvaluation> {
961 if a.shape.len() != 2 || b.shape.len() != 2 {
962 return Err(setxor_internal_error(
963 "setxor: 'rows' option requires 2-D string arrays",
964 ));
965 }
966 if a.shape[1] != b.shape[1] {
967 return Err(setxor_error(&SETXOR_ERROR_ROWS_COLUMN_MISMATCH));
968 }
969 let rows_a = a.shape[0];
970 let rows_b = b.shape[0];
971 let cols = a.shape[1];
972 let mut entries = Vec::<SymEntry<Vec<String>>>::new();
973 let mut map: HashMap<RowStringKey, usize> = HashMap::new();
974 let mut order_counter = 0usize;
975 for row in 0..rows_a {
976 let values = string_row(&a, row, cols);
977 add_sym_entry(
978 &mut entries,
979 &mut map,
980 RowStringKey(values.clone()),
981 values,
982 Origin::A,
983 row,
984 &mut order_counter,
985 );
986 }
987 for row in 0..rows_b {
988 let values = string_row(&b, row, cols);
989 add_sym_entry(
990 &mut entries,
991 &mut map,
992 RowStringKey(values.clone()),
993 values,
994 Origin::B,
995 row,
996 &mut order_counter,
997 );
998 }
999 assemble_string_rows(entries, opts, cols)
1000}
1001
1002fn add_sym_entry<K, T>(
1003 entries: &mut Vec<SymEntry<T>>,
1004 map: &mut HashMap<K, usize>,
1005 key: K,
1006 value: T,
1007 origin: Origin,
1008 index: usize,
1009 order_counter: &mut usize,
1010) where
1011 K: Eq + std::hash::Hash,
1012{
1013 match map.entry(key) {
1014 Entry::Occupied(occ) => {
1015 let entry = &mut entries[*occ.get()];
1016 match origin {
1017 Origin::A => {
1018 if entry.a_index.is_none() {
1019 entry.a_index = Some(index);
1020 }
1021 }
1022 Origin::B => {
1023 if entry.b_index.is_none() {
1024 entry.b_index = Some(index);
1025 }
1026 }
1027 }
1028 }
1029 Entry::Vacant(v) => {
1030 let entry_idx = entries.len();
1031 let (a_index, b_index) = match origin {
1032 Origin::A => (Some(index), None),
1033 Origin::B => (None, Some(index)),
1034 };
1035 entries.push(SymEntry {
1036 value,
1037 a_index,
1038 b_index,
1039 order_rank: *order_counter,
1040 });
1041 v.insert(entry_idx);
1042 *order_counter += 1;
1043 }
1044 }
1045}
1046
1047fn symmetric_order<T>(
1048 entries: &[SymEntry<T>],
1049 opts: &SetxorOptions,
1050 compare: impl Fn(&T, &T) -> Ordering,
1051) -> Vec<usize> {
1052 let mut order = entries
1053 .iter()
1054 .enumerate()
1055 .filter_map(|(idx, entry)| {
1056 if entry.a_index.is_some() ^ entry.b_index.is_some() {
1057 Some(idx)
1058 } else {
1059 None
1060 }
1061 })
1062 .collect::<Vec<_>>();
1063 match opts.order {
1064 SetxorOrder::Sorted => {
1065 order.sort_by(|&lhs, &rhs| compare(&entries[lhs].value, &entries[rhs].value))
1066 }
1067 SetxorOrder::Stable => order.sort_by_key(|&idx| entries[idx].order_rank),
1068 }
1069 order
1070}
1071
1072fn collect_indices<T>(entries: &[SymEntry<T>], order: &[usize]) -> (Vec<f64>, Vec<f64>) {
1073 let mut ia = Vec::new();
1074 let mut ib = Vec::new();
1075 for &idx in order {
1076 let entry = &entries[idx];
1077 if let Some(a_idx) = entry.a_index {
1078 ia.push((a_idx + 1) as f64);
1079 } else if let Some(b_idx) = entry.b_index {
1080 ib.push((b_idx + 1) as f64);
1081 }
1082 }
1083 (ia, ib)
1084}
1085
1086fn index_tensors(ia: Vec<f64>, ib: Vec<f64>) -> crate::BuiltinResult<(Tensor, Tensor)> {
1087 let ia_len = ia.len();
1088 let ib_len = ib.len();
1089 let ia_tensor = Tensor::new(ia, vec![ia_len, 1])
1090 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1091 let ib_tensor = Tensor::new(ib, vec![ib_len, 1])
1092 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1093 Ok((ia_tensor, ib_tensor))
1094}
1095
1096fn is_row_vector_shape(shape: &[usize]) -> bool {
1097 match shape {
1098 [] => false,
1099 [_] => true,
1100 [rows, ..] if *rows != 1 => false,
1101 [_, _, rest @ ..] => rest.iter().all(|&dim| dim == 1),
1102 }
1103}
1104
1105fn element_row_output(a_shape: &[usize], b_shape: &[usize]) -> bool {
1106 is_row_vector_shape(a_shape) && is_row_vector_shape(b_shape)
1107}
1108
1109fn element_shape(row_output: bool, len: usize) -> Vec<usize> {
1110 if row_output {
1111 vec![1, len]
1112 } else {
1113 vec![len, 1]
1114 }
1115}
1116
1117fn numeric_output_dtype(
1118 a_dtype: NumericDType,
1119 b_dtype: NumericDType,
1120) -> crate::BuiltinResult<NumericDType> {
1121 match (a_dtype, b_dtype) {
1122 (lhs, rhs) if lhs == rhs => Ok(lhs),
1123 (NumericDType::F64, rhs) => Ok(rhs),
1124 (lhs, NumericDType::F64) => Ok(lhs),
1125 _ => Err(setxor_error(&SETXOR_ERROR_NUMERIC_CLASS_MISMATCH)),
1126 }
1127}
1128
1129fn element_meta(
1130 a_shape: &[usize],
1131 a_dtype: NumericDType,
1132 b_shape: &[usize],
1133 b_dtype: NumericDType,
1134) -> crate::BuiltinResult<ElementMeta> {
1135 Ok(ElementMeta {
1136 row_output: element_row_output(a_shape, b_shape),
1137 dtype: numeric_output_dtype(a_dtype, b_dtype)?,
1138 })
1139}
1140
1141fn assemble_numeric(
1142 entries: Vec<SymEntry<f64>>,
1143 opts: &SetxorOptions,
1144 meta: &ElementMeta,
1145) -> crate::BuiltinResult<SetxorEvaluation> {
1146 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_f64(*lhs, *rhs));
1147 let values = order
1148 .iter()
1149 .map(|&idx| entries[idx].value)
1150 .collect::<Vec<_>>();
1151 let (ia, ib) = collect_indices(&entries, &order);
1152 let value_tensor = Tensor::new_with_dtype(
1153 values,
1154 element_shape(meta.row_output, order.len()),
1155 meta.dtype,
1156 )
1157 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1158 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1159 Ok(SetxorEvaluation::new(
1160 tensor::tensor_into_value(value_tensor),
1161 ia_tensor,
1162 ib_tensor,
1163 ))
1164}
1165
1166fn assemble_numeric_rows(
1167 entries: Vec<SymEntry<Vec<f64>>>,
1168 opts: &SetxorOptions,
1169 cols: usize,
1170 dtype: NumericDType,
1171) -> crate::BuiltinResult<SetxorEvaluation> {
1172 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_numeric_rows(lhs, rhs));
1173 let rows = order.len();
1174 let mut values = vec![0.0; rows * cols];
1175 for (row_pos, &entry_idx) in order.iter().enumerate() {
1176 for col in 0..cols {
1177 values[row_pos + col * rows] = entries[entry_idx].value[col];
1178 }
1179 }
1180 let (ia, ib) = collect_indices(&entries, &order);
1181 let value_tensor = Tensor::new_with_dtype(values, vec![rows, cols], dtype)
1182 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1183 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1184 Ok(SetxorEvaluation::new(
1185 tensor::tensor_into_value(value_tensor),
1186 ia_tensor,
1187 ib_tensor,
1188 ))
1189}
1190
1191fn assemble_complex(
1192 entries: Vec<SymEntry<(f64, f64)>>,
1193 opts: &SetxorOptions,
1194 row_output: bool,
1195) -> crate::BuiltinResult<SetxorEvaluation> {
1196 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_complex(*lhs, *rhs));
1197 let values = order
1198 .iter()
1199 .map(|&idx| entries[idx].value)
1200 .collect::<Vec<_>>();
1201 let (ia, ib) = collect_indices(&entries, &order);
1202 let value_tensor = ComplexTensor::new(values, element_shape(row_output, order.len()))
1203 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1204 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1205 Ok(SetxorEvaluation::new(
1206 complex_tensor_into_value(value_tensor),
1207 ia_tensor,
1208 ib_tensor,
1209 ))
1210}
1211
1212fn assemble_complex_rows(
1213 entries: Vec<SymEntry<Vec<(f64, f64)>>>,
1214 opts: &SetxorOptions,
1215 cols: usize,
1216) -> crate::BuiltinResult<SetxorEvaluation> {
1217 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_complex_rows(lhs, rhs));
1218 let rows = order.len();
1219 let mut values = vec![(0.0, 0.0); rows * cols];
1220 for (row_pos, &entry_idx) in order.iter().enumerate() {
1221 for col in 0..cols {
1222 values[row_pos + col * rows] = entries[entry_idx].value[col];
1223 }
1224 }
1225 let (ia, ib) = collect_indices(&entries, &order);
1226 let value_tensor = ComplexTensor::new(values, vec![rows, cols])
1227 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1228 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1229 Ok(SetxorEvaluation::new(
1230 complex_tensor_into_value(value_tensor),
1231 ia_tensor,
1232 ib_tensor,
1233 ))
1234}
1235
1236fn assemble_char(
1237 entries: Vec<SymEntry<char>>,
1238 opts: &SetxorOptions,
1239 row_output: bool,
1240) -> crate::BuiltinResult<SetxorEvaluation> {
1241 let order = symmetric_order(&entries, opts, |lhs, rhs| lhs.cmp(rhs));
1242 let values = order
1243 .iter()
1244 .map(|&idx| entries[idx].value)
1245 .collect::<Vec<_>>();
1246 let (ia, ib) = collect_indices(&entries, &order);
1247 let (rows, cols) = if row_output {
1248 (1, order.len())
1249 } else {
1250 (order.len(), 1)
1251 };
1252 let value_array = CharArray::new(values, rows, cols)
1253 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1254 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1255 Ok(SetxorEvaluation::new(
1256 Value::CharArray(value_array),
1257 ia_tensor,
1258 ib_tensor,
1259 ))
1260}
1261
1262fn assemble_char_rows(
1263 entries: Vec<SymEntry<Vec<char>>>,
1264 opts: &SetxorOptions,
1265 cols: usize,
1266) -> crate::BuiltinResult<SetxorEvaluation> {
1267 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_char_rows(lhs, rhs));
1268 let rows = order.len();
1269 let mut values = vec!['\0'; rows * cols];
1270 for (row_pos, &entry_idx) in order.iter().enumerate() {
1271 for col in 0..cols {
1272 values[row_pos * cols + col] = entries[entry_idx].value[col];
1273 }
1274 }
1275 let (ia, ib) = collect_indices(&entries, &order);
1276 let value_array = CharArray::new(values, rows, cols)
1277 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1278 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1279 Ok(SetxorEvaluation::new(
1280 Value::CharArray(value_array),
1281 ia_tensor,
1282 ib_tensor,
1283 ))
1284}
1285
1286fn assemble_string(
1287 entries: Vec<SymEntry<String>>,
1288 opts: &SetxorOptions,
1289 row_output: bool,
1290) -> crate::BuiltinResult<SetxorEvaluation> {
1291 let order = symmetric_order(&entries, opts, |lhs, rhs| lhs.cmp(rhs));
1292 let values = order
1293 .iter()
1294 .map(|&idx| entries[idx].value.clone())
1295 .collect::<Vec<_>>();
1296 let (ia, ib) = collect_indices(&entries, &order);
1297 let value_array = StringArray::new(values, element_shape(row_output, order.len()))
1298 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1299 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1300 Ok(SetxorEvaluation::new(
1301 Value::StringArray(value_array),
1302 ia_tensor,
1303 ib_tensor,
1304 ))
1305}
1306
1307fn assemble_string_rows(
1308 entries: Vec<SymEntry<Vec<String>>>,
1309 opts: &SetxorOptions,
1310 cols: usize,
1311) -> crate::BuiltinResult<SetxorEvaluation> {
1312 let order = symmetric_order(&entries, opts, |lhs, rhs| compare_string_rows(lhs, rhs));
1313 let rows = order.len();
1314 let mut values = vec![String::new(); rows * cols];
1315 for (row_pos, &entry_idx) in order.iter().enumerate() {
1316 for col in 0..cols {
1317 values[row_pos + col * rows] = entries[entry_idx].value[col].clone();
1318 }
1319 }
1320 let (ia, ib) = collect_indices(&entries, &order);
1321 let value_array = StringArray::new(values, vec![rows, cols])
1322 .map_err(|e| setxor_internal_error(format!("setxor: {e}")))?;
1323 let (ia_tensor, ib_tensor) = index_tensors(ia, ib)?;
1324 Ok(SetxorEvaluation::new(
1325 Value::StringArray(value_array),
1326 ia_tensor,
1327 ib_tensor,
1328 ))
1329}
1330
1331#[derive(Debug, Clone)]
1332pub struct SetxorEvaluation {
1333 values: Value,
1334 ia: Tensor,
1335 ib: Tensor,
1336}
1337
1338impl SetxorEvaluation {
1339 fn new(values: Value, ia: Tensor, ib: Tensor) -> Self {
1340 Self { values, ia, ib }
1341 }
1342
1343 pub fn into_values_value(self) -> Value {
1344 self.values
1345 }
1346
1347 pub fn into_triple(self) -> (Value, Value, Value) {
1348 (
1349 self.values,
1350 tensor::tensor_into_value(self.ia),
1351 tensor::tensor_into_value(self.ib),
1352 )
1353 }
1354
1355 pub fn values_value(&self) -> Value {
1356 self.values.clone()
1357 }
1358
1359 pub fn ia_value(&self) -> Value {
1360 tensor::tensor_into_value(self.ia.clone())
1361 }
1362
1363 pub fn ib_value(&self) -> Value {
1364 tensor::tensor_into_value(self.ib.clone())
1365 }
1366}
1367
1368fn numeric_key(value: f64, origin: Origin, index: usize) -> NumericKey {
1369 if value.is_nan() {
1370 NumericKey::UniqueNan(origin, index)
1371 } else {
1372 NumericKey::Value(canonicalize_f64(value))
1373 }
1374}
1375
1376fn numeric_row_key(values: &[f64], origin: Origin, row: usize) -> NumericRowKey {
1377 if values.iter().any(|value| value.is_nan()) {
1378 NumericRowKey::UniqueNan(origin, row)
1379 } else {
1380 NumericRowKey::Values(
1381 values
1382 .iter()
1383 .map(|&value| canonicalize_f64(value))
1384 .collect(),
1385 )
1386 }
1387}
1388
1389fn complex_element_key(value: (f64, f64), origin: Origin, index: usize) -> ComplexElementKey {
1390 if complex_is_nan(value) {
1391 ComplexElementKey::UniqueNan(origin, index)
1392 } else {
1393 ComplexElementKey::Value(ComplexKey::new(value))
1394 }
1395}
1396
1397fn complex_row_key(values: &[(f64, f64)], origin: Origin, row: usize) -> ComplexRowKey {
1398 if values.iter().any(|&value| complex_is_nan(value)) {
1399 ComplexRowKey::UniqueNan(origin, row)
1400 } else {
1401 ComplexRowKey::Values(values.iter().map(|&value| ComplexKey::new(value)).collect())
1402 }
1403}
1404
1405fn numeric_row(tensor: &Tensor, row: usize, cols: usize) -> Vec<f64> {
1406 (0..cols)
1407 .map(|col| tensor.data[row + col * tensor.shape[0]])
1408 .collect()
1409}
1410
1411fn complex_row(tensor: &ComplexTensor, row: usize, cols: usize) -> Vec<(f64, f64)> {
1412 (0..cols)
1413 .map(|col| tensor.data[row + col * tensor.shape[0]])
1414 .collect()
1415}
1416
1417fn char_row(array: &CharArray, row: usize) -> Vec<char> {
1418 (0..array.cols)
1419 .map(|col| array.data[row * array.cols + col])
1420 .collect()
1421}
1422
1423fn string_row(array: &StringArray, row: usize, cols: usize) -> Vec<String> {
1424 (0..cols)
1425 .map(|col| array.data[row + col * array.shape[0]].clone())
1426 .collect()
1427}
1428
1429fn canonicalize_f64(value: f64) -> u64 {
1430 if value == 0.0 {
1431 0
1432 } else {
1433 value.to_bits()
1434 }
1435}
1436
1437fn compare_f64(a: f64, b: f64) -> Ordering {
1438 if a.is_nan() {
1439 if b.is_nan() {
1440 Ordering::Equal
1441 } else {
1442 Ordering::Greater
1443 }
1444 } else if b.is_nan() {
1445 Ordering::Less
1446 } else {
1447 a.partial_cmp(&b).unwrap_or(Ordering::Equal)
1448 }
1449}
1450
1451fn compare_numeric_rows(a: &[f64], b: &[f64]) -> Ordering {
1452 for (lhs, rhs) in a.iter().zip(b.iter()) {
1453 let ord = compare_f64(*lhs, *rhs);
1454 if ord != Ordering::Equal {
1455 return ord;
1456 }
1457 }
1458 Ordering::Equal
1459}
1460
1461impl ComplexKey {
1462 fn new(value: (f64, f64)) -> Self {
1463 Self {
1464 re: canonicalize_f64(value.0),
1465 im: canonicalize_f64(value.1),
1466 }
1467 }
1468}
1469
1470fn complex_is_nan(value: (f64, f64)) -> bool {
1471 value.0.is_nan() || value.1.is_nan()
1472}
1473
1474fn compare_complex(a: (f64, f64), b: (f64, f64)) -> Ordering {
1475 match (complex_is_nan(a), complex_is_nan(b)) {
1476 (true, true) => Ordering::Equal,
1477 (true, false) => Ordering::Greater,
1478 (false, true) => Ordering::Less,
1479 (false, false) => {
1480 let mag_cmp = compare_f64(a.0.hypot(a.1), b.0.hypot(b.1));
1481 if mag_cmp != Ordering::Equal {
1482 return mag_cmp;
1483 }
1484 let phase_cmp = compare_f64(a.1.atan2(a.0), b.1.atan2(b.0));
1485 if phase_cmp != Ordering::Equal {
1486 return phase_cmp;
1487 }
1488 let re_cmp = compare_f64(a.0, b.0);
1489 if re_cmp != Ordering::Equal {
1490 re_cmp
1491 } else {
1492 compare_f64(a.1, b.1)
1493 }
1494 }
1495 }
1496}
1497
1498fn compare_complex_rows(a: &[(f64, f64)], b: &[(f64, f64)]) -> Ordering {
1499 for (lhs, rhs) in a.iter().zip(b.iter()) {
1500 let ord = compare_complex(*lhs, *rhs);
1501 if ord != Ordering::Equal {
1502 return ord;
1503 }
1504 }
1505 Ordering::Equal
1506}
1507
1508fn compare_char_rows(a: &[char], b: &[char]) -> Ordering {
1509 for (lhs, rhs) in a.iter().zip(b.iter()) {
1510 let ord = lhs.cmp(rhs);
1511 if ord != Ordering::Equal {
1512 return ord;
1513 }
1514 }
1515 Ordering::Equal
1516}
1517
1518fn compare_string_rows(a: &[String], b: &[String]) -> Ordering {
1519 for (lhs, rhs) in a.iter().zip(b.iter()) {
1520 let ord = lhs.cmp(rhs);
1521 if ord != Ordering::Equal {
1522 return ord;
1523 }
1524 }
1525 Ordering::Equal
1526}
1527
1528#[cfg(test)]
1529mod tests {
1530 use super::*;
1531 use crate::builtins::common::test_support;
1532 use runmat_accelerate_api::HostTensorView;
1533 use runmat_builtins::{IntValue, ResolveContext, Type};
1534
1535 fn evaluate_sync(a: Value, b: Value, rest: &[Value]) -> crate::BuiltinResult<SetxorEvaluation> {
1536 futures::executor::block_on(evaluate(a, b, rest))
1537 }
1538
1539 fn builtin_sync(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1540 futures::executor::block_on(setxor_builtin(a, b, rest))
1541 }
1542
1543 #[test]
1544 fn setxor_type_resolver_numeric() {
1545 assert_eq!(
1546 set_values_output_type(
1547 &[Type::tensor(), Type::tensor()],
1548 &ResolveContext::new(Vec::new()),
1549 ),
1550 Type::tensor()
1551 );
1552 }
1553
1554 #[test]
1555 fn setxor_numeric_sorted_default_with_indices() {
1556 let a = Tensor::new(vec![5.0, 1.0, 3.0, 3.0, 3.0], vec![5, 1]).unwrap();
1557 let b = Tensor::new(vec![4.0, 1.0, 2.0], vec![3, 1]).unwrap();
1558 let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("setxor");
1559 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1560 assert_eq!(values.data, vec![2.0, 3.0, 4.0, 5.0]);
1561 assert_eq!(values.shape, vec![4, 1]);
1562 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1563 assert_eq!(ia.data, vec![3.0, 1.0]);
1564 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1565 assert_eq!(ib.data, vec![3.0, 1.0]);
1566 }
1567
1568 #[test]
1569 fn setxor_numeric_preserves_row_vector_shape_when_both_inputs_are_rows() {
1570 let a = Tensor::new(vec![5.0, 1.0, 3.0], vec![1, 3]).unwrap();
1571 let b = Tensor::new(vec![4.0, 1.0, 2.0], vec![1, 3]).unwrap();
1572 let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("setxor");
1573 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1574 assert_eq!(values.data, vec![2.0, 3.0, 4.0, 5.0]);
1575 assert_eq!(values.shape, vec![1, 4]);
1576 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1577 assert_eq!(ia.shape, vec![2, 1]);
1578 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1579 assert_eq!(ib.shape, vec![2, 1]);
1580 }
1581
1582 #[test]
1583 fn setxor_numeric_preserves_matching_dtype() {
1584 let a = Tensor::new_with_dtype(vec![5.0, 1.0, 3.0], vec![1, 3], NumericDType::U32).unwrap();
1585 let b = Tensor::new_with_dtype(vec![5.0, 2.0], vec![1, 2], NumericDType::U32).unwrap();
1586 let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("setxor");
1587 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1588 assert_eq!(values.data, vec![1.0, 2.0, 3.0]);
1589 assert_eq!(values.shape, vec![1, 3]);
1590 assert_eq!(values.dtype, NumericDType::U32);
1591 }
1592
1593 #[test]
1594 fn setxor_numeric_double_and_nondouble_returns_nondouble_dtype() {
1595 let a = Tensor::new_with_dtype(vec![5.0, 1.0, 3.0], vec![1, 3], NumericDType::U32).unwrap();
1596 let b = Tensor::new(vec![5.0, 2.0], vec![1, 2]).unwrap();
1597 let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("setxor");
1598 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1599 assert_eq!(values.data, vec![1.0, 2.0, 3.0]);
1600 assert_eq!(values.dtype, NumericDType::U32);
1601 }
1602
1603 #[test]
1604 fn setxor_numeric_rejects_incompatible_nondouble_classes() {
1605 let a = Tensor::new_with_dtype(vec![1.0, 2.0], vec![1, 2], NumericDType::U8).unwrap();
1606 let b = Tensor::new_with_dtype(vec![2.0, 3.0], vec![1, 2], NumericDType::U32).unwrap();
1607 let err = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).unwrap_err();
1608 assert_eq!(
1609 err.identifier(),
1610 SETXOR_ERROR_NUMERIC_CLASS_MISMATCH.identifier
1611 );
1612 }
1613
1614 #[test]
1615 fn setxor_numeric_stable_order() {
1616 let a = Tensor::new(vec![5.0, 1.0, 3.0, 3.0, 3.0], vec![5, 1]).unwrap();
1617 let b = Tensor::new(vec![4.0, 1.0, 2.0], vec![3, 1]).unwrap();
1618 let eval =
1619 evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("stable")]).unwrap();
1620 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1621 assert_eq!(values.data, vec![5.0, 3.0, 4.0, 2.0]);
1622 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1623 assert_eq!(ia.data, vec![1.0, 3.0]);
1624 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1625 assert_eq!(ib.data, vec![1.0, 3.0]);
1626 }
1627
1628 #[test]
1629 fn setxor_treats_nan_values_as_distinct() {
1630 let a = Tensor::new(vec![5.0, f64::NAN, f64::NAN], vec![3, 1]).unwrap();
1631 let b = Tensor::new(vec![5.0, f64::NAN, f64::NAN], vec![3, 1]).unwrap();
1632 let eval = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[]).expect("setxor");
1633 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1634 assert_eq!(values.shape, vec![4, 1]);
1635 assert!(values.data.iter().all(|value| value.is_nan()));
1636 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1637 assert_eq!(ia.data, vec![2.0, 3.0]);
1638 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1639 assert_eq!(ib.data, vec![2.0, 3.0]);
1640 }
1641
1642 #[test]
1643 fn setxor_numeric_rows_sorted() {
1644 let a = Tensor::new(
1645 vec![
1646 7.0, 7.0, 7.0, 1.0, 4.0, 8.0, 7.0, 7.0, 2.0, 5.0, 9.0, 1.0, 1.0, 3.0, 6.0,
1647 ],
1648 vec![5, 3],
1649 )
1650 .unwrap();
1651 let b = Tensor::new(
1652 vec![1.0, 4.0, 7.0, 2.0, 5.0, 7.0, 3.0, 6.0, 2.0],
1653 vec![3, 3],
1654 )
1655 .unwrap();
1656 let eval =
1657 evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")]).unwrap();
1658 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1659 assert_eq!(values.shape, vec![3, 3]);
1660 assert_eq!(
1661 values.data,
1662 vec![7.0, 7.0, 7.0, 7.0, 7.0, 8.0, 1.0, 2.0, 9.0]
1663 );
1664 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1665 assert_eq!(ia.data, vec![2.0, 1.0]);
1666 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1667 assert_eq!(ib.data, vec![3.0]);
1668 }
1669
1670 #[test]
1671 fn setxor_complex_values() {
1672 let a = ComplexTensor::new(vec![(1.0, 1.0), (2.0, 0.0)], vec![2, 1]).unwrap();
1673 let b = ComplexTensor::new(vec![(2.0, 0.0), (3.0, 0.0)], vec![2, 1]).unwrap();
1674 let eval =
1675 evaluate_sync(Value::ComplexTensor(a), Value::ComplexTensor(b), &[]).expect("setxor");
1676 let Value::ComplexTensor(values) = eval.values_value() else {
1677 panic!("expected complex tensor");
1678 };
1679 assert_eq!(values.data, vec![(1.0, 1.0), (3.0, 0.0)]);
1680 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1681 assert_eq!(ia.data, vec![1.0]);
1682 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1683 assert_eq!(ib.data, vec![2.0]);
1684 }
1685
1686 #[test]
1687 fn setxor_promotes_real_input_to_complex_domain() {
1688 let a = ComplexTensor::new(vec![(1.0, 1.0), (2.0, 0.0)], vec![1, 2]).unwrap();
1689 let b = Tensor::new(vec![2.0, 3.0], vec![1, 2]).unwrap();
1690 let eval = evaluate_sync(Value::ComplexTensor(a), Value::Tensor(b), &[]).expect("setxor");
1691 let Value::ComplexTensor(values) = eval.values_value() else {
1692 panic!("expected complex tensor");
1693 };
1694 assert_eq!(values.data, vec![(1.0, 1.0), (3.0, 0.0)]);
1695 assert_eq!(values.shape, vec![1, 2]);
1696 }
1697
1698 #[test]
1699 fn setxor_complex_sorted_uses_phase_after_magnitude() {
1700 let a = ComplexTensor::new(vec![(0.0, 1.0)], vec![1, 1]).unwrap();
1701 let b = ComplexTensor::new(vec![(1.0, 0.0)], vec![1, 1]).unwrap();
1702 let eval =
1703 evaluate_sync(Value::ComplexTensor(a), Value::ComplexTensor(b), &[]).expect("setxor");
1704 let Value::ComplexTensor(values) = eval.values_value() else {
1705 panic!("expected complex tensor");
1706 };
1707 assert_eq!(values.data, vec![(1.0, 0.0), (0.0, 1.0)]);
1708 assert_eq!(values.shape, vec![1, 2]);
1709 }
1710
1711 #[test]
1712 fn setxor_char_elements() {
1713 let a = CharArray::new(vec!['d', 'o', 'g'], 1, 3).unwrap();
1714 let b = CharArray::new(vec!['d', 'i', 'g'], 1, 3).unwrap();
1715 let eval = evaluate_sync(Value::CharArray(a), Value::CharArray(b), &[]).expect("setxor");
1716 let Value::CharArray(values) = eval.values_value() else {
1717 panic!("expected char array");
1718 };
1719 assert_eq!(values.data, vec!['i', 'o']);
1720 assert_eq!((values.rows, values.cols), (1, 2));
1721 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1722 assert_eq!(ia.data, vec![2.0]);
1723 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1724 assert_eq!(ib.data, vec![2.0]);
1725 }
1726
1727 #[test]
1728 fn setxor_char_and_numeric_compare_character_codes() {
1729 let a = CharArray::new_row("abc");
1730 let b = Tensor::new(vec![98.0, 100.0], vec![1, 2]).unwrap();
1731 let eval = evaluate_sync(Value::CharArray(a), Value::Tensor(b), &[]).expect("setxor");
1732 let Value::CharArray(values) = eval.values_value() else {
1733 panic!("expected char array");
1734 };
1735 assert_eq!(values.data, vec!['a', 'c', 'd']);
1736 assert_eq!((values.rows, values.cols), (1, 3));
1737 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1738 assert_eq!(ia.data, vec![1.0, 3.0]);
1739 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1740 assert_eq!(ib.data, vec![2.0]);
1741 }
1742
1743 #[test]
1744 fn setxor_string_and_char_vector_compare_strings() {
1745 let a =
1746 StringArray::new(vec!["alpha".to_string(), "beta".to_string()], vec![1, 2]).unwrap();
1747 let b = CharArray::new_row("beta");
1748 let eval = evaluate_sync(Value::StringArray(a), Value::CharArray(b), &[]).expect("setxor");
1749 let Value::StringArray(values) = eval.values_value() else {
1750 panic!("expected string array");
1751 };
1752 assert_eq!(values.data, vec!["alpha".to_string()]);
1753 assert_eq!(values.shape, vec![1, 1]);
1754 }
1755
1756 #[test]
1757 fn setxor_string_rows_stable() {
1758 let a = StringArray::new(
1759 vec![
1760 "alpha".to_string(),
1761 "gamma".to_string(),
1762 "beta".to_string(),
1763 "beta".to_string(),
1764 ],
1765 vec![2, 2],
1766 )
1767 .unwrap();
1768 let b = StringArray::new(
1769 vec![
1770 "gamma".to_string(),
1771 "delta".to_string(),
1772 "beta".to_string(),
1773 "beta".to_string(),
1774 ],
1775 vec![2, 2],
1776 )
1777 .unwrap();
1778 let eval = evaluate_sync(
1779 Value::StringArray(a),
1780 Value::StringArray(b),
1781 &[Value::from("rows"), Value::from("stable")],
1782 )
1783 .unwrap();
1784 let Value::StringArray(values) = eval.values_value() else {
1785 panic!("expected string array");
1786 };
1787 assert_eq!(values.shape, vec![2, 2]);
1788 assert_eq!(
1789 values.data,
1790 vec![
1791 "alpha".to_string(),
1792 "delta".to_string(),
1793 "beta".to_string(),
1794 "beta".to_string()
1795 ]
1796 );
1797 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1798 assert_eq!(ia.data, vec![1.0]);
1799 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1800 assert_eq!(ib.data, vec![2.0]);
1801 }
1802
1803 #[test]
1804 fn setxor_gpu_roundtrip() {
1805 test_support::with_test_provider(|provider| {
1806 let a = Tensor::new(vec![4.0, 1.0, 2.0], vec![3, 1]).unwrap();
1807 let b = Tensor::new(vec![2.0, 5.0], vec![2, 1]).unwrap();
1808 let view_a = HostTensorView {
1809 data: &a.data,
1810 shape: &a.shape,
1811 };
1812 let view_b = HostTensorView {
1813 data: &b.data,
1814 shape: &b.shape,
1815 };
1816 let handle_a = provider.upload(&view_a).expect("upload A");
1817 let handle_b = provider.upload(&view_b).expect("upload B");
1818 let eval = evaluate_sync(
1819 Value::GpuTensor(handle_a),
1820 Value::GpuTensor(handle_b),
1821 &[Value::from("stable")],
1822 )
1823 .expect("setxor");
1824 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1825 assert_eq!(values.data, vec![4.0, 1.0, 5.0]);
1826 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1827 assert_eq!(ia.data, vec![1.0, 2.0]);
1828 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1829 assert_eq!(ib.data, vec![2.0]);
1830 });
1831 }
1832
1833 #[test]
1834 fn setxor_gpu_real_and_host_complex_match_host_promotion() {
1835 test_support::with_test_provider(|provider| {
1836 let a = Tensor::new(vec![2.0, 3.0], vec![1, 2]).unwrap();
1837 let view_a = HostTensorView {
1838 data: &a.data,
1839 shape: &a.shape,
1840 };
1841 let handle_a = provider.upload(&view_a).expect("upload A");
1842 let b = ComplexTensor::new(vec![(1.0, 1.0), (2.0, 0.0)], vec![1, 2]).unwrap();
1843 let eval = evaluate_sync(Value::GpuTensor(handle_a), Value::ComplexTensor(b), &[])
1844 .expect("setxor");
1845 let Value::ComplexTensor(values) = eval.values_value() else {
1846 panic!("expected complex tensor");
1847 };
1848 assert_eq!(values.data, vec![(1.0, 1.0), (3.0, 0.0)]);
1849 assert_eq!(values.shape, vec![1, 2]);
1850 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1851 assert_eq!(ia.data, vec![2.0]);
1852 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1853 assert_eq!(ib.data, vec![1.0]);
1854 });
1855 }
1856
1857 #[test]
1858 fn setxor_rejects_legacy_option() {
1859 let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
1860 let err = evaluate_sync(
1861 Value::Tensor(tensor.clone()),
1862 Value::Tensor(tensor),
1863 &[Value::from("legacy")],
1864 )
1865 .unwrap_err();
1866 assert_eq!(
1867 err.identifier(),
1868 SETXOR_ERROR_LEGACY_OPTION_UNSUPPORTED.identifier
1869 );
1870 }
1871
1872 #[test]
1873 fn setxor_rejects_conflicting_order_options() {
1874 let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
1875 let err = evaluate_sync(
1876 Value::Tensor(tensor.clone()),
1877 Value::Tensor(tensor),
1878 &[Value::from("stable"), Value::from("sorted")],
1879 )
1880 .unwrap_err();
1881 assert_eq!(
1882 err.identifier(),
1883 SETXOR_ERROR_CONFLICTING_ORDER_OPTIONS.identifier
1884 );
1885 }
1886
1887 #[test]
1888 fn setxor_rows_dimension_mismatch() {
1889 let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
1890 let b = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
1891 let err =
1892 evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")]).unwrap_err();
1893 assert_eq!(
1894 err.identifier(),
1895 SETXOR_ERROR_ROWS_COLUMN_MISMATCH.identifier
1896 );
1897 }
1898
1899 #[test]
1900 fn setxor_accepts_scalar_inputs() {
1901 let eval =
1902 evaluate_sync(Value::Int(IntValue::I32(1)), Value::Num(3.0), &[]).expect("setxor");
1903 let values = tensor::value_into_tensor_for("setxor", eval.values_value()).unwrap();
1904 assert_eq!(values.data, vec![1.0, 3.0]);
1905 let ia = tensor::value_into_tensor_for("setxor", eval.ia_value()).unwrap();
1906 assert_eq!(ia.data, vec![1.0]);
1907 let ib = tensor::value_into_tensor_for("setxor", eval.ib_value()).unwrap();
1908 assert_eq!(ib.data, vec![1.0]);
1909 }
1910
1911 #[test]
1912 fn setxor_rejects_more_than_three_outputs() {
1913 let _guard = crate::output_count::push_output_count(Some(4));
1914 let tensor = Tensor::new(vec![1.0, 2.0], vec![1, 2]).unwrap();
1915 let err = builtin_sync(
1916 Value::Tensor(tensor.clone()),
1917 Value::Tensor(tensor),
1918 Vec::new(),
1919 )
1920 .expect_err("too many outputs should fail");
1921 assert_eq!(err.identifier(), SETXOR_ERROR_TOO_MANY_OUTPUTS.identifier);
1922 }
1923}