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