1use std::cmp::Ordering;
4
5use runmat_accelerate_api::{
6 GpuTensorHandle, SortComparison as ProviderSortComparison, SortOrder as ProviderSortOrder,
7};
8use runmat_builtins::{
9 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
10 BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
11 BuiltinIntegerInputAvailability, BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule,
12 BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule,
13 BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
14 BuiltinSignatureDescriptor,
15};
16use runmat_macros::runtime_builtin;
17use runmat_value::{
18 ComplexStorage, ComplexTensor, IntValue, IntegerStorage, LogicalArray, NumericStorage, Tensor,
19 Value,
20};
21
22use super::{float_order::SetFloat, integer_order, type_resolvers::tensor_output_type};
23use crate::build_runtime_error;
24use crate::builtins::common::arg_tokens::{tokens_from_values, ArgToken};
25use crate::builtins::common::gpu_helpers;
26use crate::builtins::common::spec::{
27 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
28 ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
29};
30use crate::builtins::common::tensor;
31
32#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::array::sorting_sets::sort")]
33pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
34 name: "sort",
35 op_kind: GpuOpKind::Custom("sort"),
36 supported_precisions: &[ScalarType::F32, ScalarType::F64],
37 broadcast: BroadcastSemantics::None,
38 provider_hooks: &[ProviderHook::Custom("sort_dim")],
39 constant_strategy: ConstantStrategy::InlineLiteral,
40 residency: ResidencyPolicy::NewHandle,
41 nan_mode: ReductionNaN::Include,
42 two_pass_threshold: None,
43 workgroup_size: None,
44 accepts_nan_mode: true,
45 notes: "Plain real tensors may use the provider sort hook; typed integer, logical, complex, explicit missing-placement, and unsupported-provider paths gather through authoritative host storage and restore both outputs to the owning provider.",
46};
47
48#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::array::sorting_sets::sort")]
49pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
50 name: "sort",
51 shape: ShapeRequirements::Any,
52 constant_strategy: ConstantStrategy::InlineLiteral,
53 elementwise: None,
54 reduction: None,
55 emits_nan: true,
56 notes: "Sorting breaks fusion chains; host fallback may gather upstream tensors before restoring new resident output handles.",
57};
58
59const BUILTIN_NAME: &str = "sort";
60
61const SORT_OUTPUT_B: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
62 name: "B",
63 ty: BuiltinParamType::Any,
64 arity: BuiltinParamArity::Required,
65 default: None,
66 description: "Sorted values.",
67}];
68
69const SORT_OUTPUT_BI: [BuiltinParamDescriptor; 2] = [
70 BuiltinParamDescriptor {
71 name: "B",
72 ty: BuiltinParamType::Any,
73 arity: BuiltinParamArity::Required,
74 default: None,
75 description: "Sorted values.",
76 },
77 BuiltinParamDescriptor {
78 name: "I",
79 ty: BuiltinParamType::NumericArray,
80 arity: BuiltinParamArity::Required,
81 default: None,
82 description: "One-based index permutation for each sorted slice.",
83 },
84];
85
86const SORT_INPUTS_A: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
87 name: "A",
88 ty: BuiltinParamType::Any,
89 arity: BuiltinParamArity::Required,
90 default: None,
91 description: "Input array.",
92}];
93
94const SORT_INPUTS_A_ARG1: [BuiltinParamDescriptor; 2] = [
95 BuiltinParamDescriptor {
96 name: "A",
97 ty: BuiltinParamType::Any,
98 arity: BuiltinParamArity::Required,
99 default: None,
100 description: "Input array.",
101 },
102 BuiltinParamDescriptor {
103 name: "arg1",
104 ty: BuiltinParamType::Any,
105 arity: BuiltinParamArity::Required,
106 default: None,
107 description: "Dimension selector or direction token ('ascend'/'descend').",
108 },
109];
110
111const SORT_INPUTS_A_ARG1_ARG2: [BuiltinParamDescriptor; 3] = [
112 BuiltinParamDescriptor {
113 name: "A",
114 ty: BuiltinParamType::Any,
115 arity: BuiltinParamArity::Required,
116 default: None,
117 description: "Input array.",
118 },
119 BuiltinParamDescriptor {
120 name: "arg1",
121 ty: BuiltinParamType::Any,
122 arity: BuiltinParamArity::Required,
123 default: None,
124 description: "Dimension selector, placeholder, or direction token.",
125 },
126 BuiltinParamDescriptor {
127 name: "arg2",
128 ty: BuiltinParamType::Any,
129 arity: BuiltinParamArity::Required,
130 default: None,
131 description: "Dimension selector or direction token.",
132 },
133];
134
135const SORT_INPUTS_COMPARISON_METHOD: [BuiltinParamDescriptor; 4] = [
136 BuiltinParamDescriptor {
137 name: "A",
138 ty: BuiltinParamType::Any,
139 arity: BuiltinParamArity::Required,
140 default: None,
141 description: "Input array.",
142 },
143 BuiltinParamDescriptor {
144 name: "arg",
145 ty: BuiltinParamType::Any,
146 arity: BuiltinParamArity::Variadic,
147 default: None,
148 description: "Optional dimension/direction arguments.",
149 },
150 BuiltinParamDescriptor {
151 name: "name",
152 ty: BuiltinParamType::StringScalar,
153 arity: BuiltinParamArity::Required,
154 default: Some("\"ComparisonMethod\""),
155 description: "Name-value option key.",
156 },
157 BuiltinParamDescriptor {
158 name: "method",
159 ty: BuiltinParamType::StringScalar,
160 arity: BuiltinParamArity::Required,
161 default: Some("\"auto\""),
162 description: "Comparison method: 'auto', 'real', or 'abs'.",
163 },
164];
165
166const SORT_INPUTS_MISSING_PLACEMENT: [BuiltinParamDescriptor; 4] = [
167 BuiltinParamDescriptor {
168 name: "A",
169 ty: BuiltinParamType::Any,
170 arity: BuiltinParamArity::Required,
171 default: None,
172 description: "Input array.",
173 },
174 BuiltinParamDescriptor {
175 name: "arg",
176 ty: BuiltinParamType::Any,
177 arity: BuiltinParamArity::Variadic,
178 default: None,
179 description: "Optional dimension/direction arguments.",
180 },
181 BuiltinParamDescriptor {
182 name: "name",
183 ty: BuiltinParamType::StringScalar,
184 arity: BuiltinParamArity::Required,
185 default: Some("\"MissingPlacement\""),
186 description: "Name-value option key.",
187 },
188 BuiltinParamDescriptor {
189 name: "placement",
190 ty: BuiltinParamType::StringScalar,
191 arity: BuiltinParamArity::Required,
192 default: Some("\"auto\""),
193 description: "Missing-value placement: 'auto', 'first', or 'last'.",
194 },
195];
196
197const SORT_SIGNATURES: [BuiltinSignatureDescriptor; 10] = [
198 BuiltinSignatureDescriptor {
199 label: "B = sort(A)",
200 inputs: &SORT_INPUTS_A,
201 outputs: &SORT_OUTPUT_B,
202 },
203 BuiltinSignatureDescriptor {
204 label: "B = sort(A, arg1)",
205 inputs: &SORT_INPUTS_A_ARG1,
206 outputs: &SORT_OUTPUT_B,
207 },
208 BuiltinSignatureDescriptor {
209 label: "B = sort(A, arg1, arg2)",
210 inputs: &SORT_INPUTS_A_ARG1_ARG2,
211 outputs: &SORT_OUTPUT_B,
212 },
213 BuiltinSignatureDescriptor {
214 label: "B = sort(A, ..., \"ComparisonMethod\", method)",
215 inputs: &SORT_INPUTS_COMPARISON_METHOD,
216 outputs: &SORT_OUTPUT_B,
217 },
218 BuiltinSignatureDescriptor {
219 label: "B = sort(A, ..., \"MissingPlacement\", placement)",
220 inputs: &SORT_INPUTS_MISSING_PLACEMENT,
221 outputs: &SORT_OUTPUT_B,
222 },
223 BuiltinSignatureDescriptor {
224 label: "[B, I] = sort(A)",
225 inputs: &SORT_INPUTS_A,
226 outputs: &SORT_OUTPUT_BI,
227 },
228 BuiltinSignatureDescriptor {
229 label: "[B, I] = sort(A, arg1)",
230 inputs: &SORT_INPUTS_A_ARG1,
231 outputs: &SORT_OUTPUT_BI,
232 },
233 BuiltinSignatureDescriptor {
234 label: "[B, I] = sort(A, arg1, arg2)",
235 inputs: &SORT_INPUTS_A_ARG1_ARG2,
236 outputs: &SORT_OUTPUT_BI,
237 },
238 BuiltinSignatureDescriptor {
239 label: "[B, I] = sort(A, ..., \"ComparisonMethod\", method)",
240 inputs: &SORT_INPUTS_COMPARISON_METHOD,
241 outputs: &SORT_OUTPUT_BI,
242 },
243 BuiltinSignatureDescriptor {
244 label: "[B, I] = sort(A, ..., \"MissingPlacement\", placement)",
245 inputs: &SORT_INPUTS_MISSING_PLACEMENT,
246 outputs: &SORT_OUTPUT_BI,
247 },
248];
249
250const SORT_ERROR_INVALID_DIMENSION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
251 code: "RM.SORT.INVALID_DIMENSION",
252 identifier: Some("RunMat:sort:InvalidDimension"),
253 when: "Dimension argument is non-positive, non-integer, or otherwise invalid.",
254 message: "sort: invalid dimension argument",
255};
256
257const SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING: BuiltinErrorDescriptor =
258 BuiltinErrorDescriptor {
259 code: "RM.SORT.COMPARISON_METHOD_REQUIRES_STRING",
260 identifier: Some("RunMat:sort:ComparisonMethodRequiresString"),
261 when: "ComparisonMethod option value is not string-like.",
262 message: "sort: 'ComparisonMethod' requires a string value",
263 };
264
265const SORT_ERROR_COMPARISON_METHOD_UNKNOWN: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
266 code: "RM.SORT.COMPARISON_METHOD_UNKNOWN",
267 identifier: Some("RunMat:sort:ComparisonMethodUnknown"),
268 when: "ComparisonMethod option value is not one of 'auto'/'real'/'abs'.",
269 message: "sort: unsupported ComparisonMethod",
270};
271
272const SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING: BuiltinErrorDescriptor =
273 BuiltinErrorDescriptor {
274 code: "RM.SORT.MISSINGPLACEMENT_REQUIRES_STRING",
275 identifier: Some("RunMat:sort:MissingPlacementRequiresString"),
276 when: "MissingPlacement option value is not string-like.",
277 message: "sort: 'MissingPlacement' requires a string value",
278 };
279
280const SORT_ERROR_MISSINGPLACEMENT_UNKNOWN: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
281 code: "RM.SORT.MISSINGPLACEMENT_UNKNOWN",
282 identifier: Some("RunMat:sort:MissingPlacementUnknown"),
283 when: "MissingPlacement option value is not one of 'auto'/'first'/'last'.",
284 message: "sort: unsupported MissingPlacement",
285};
286
287const SORT_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
288 code: "RM.SORT.INVALID_ARGUMENT",
289 identifier: Some("RunMat:sort:InvalidArgument"),
290 when: "Parser encounters invalid or unrecognized option/value arguments.",
291 message: "sort: invalid argument sequence",
292};
293
294const SORT_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
295 code: "RM.SORT.INTERNAL",
296 identifier: Some("RunMat:sort:Internal"),
297 when: "Internal conversion, allocation, or provider result construction fails.",
298 message: "sort: internal operation failed",
299};
300
301const SORT_ERRORS: [BuiltinErrorDescriptor; 7] = [
302 SORT_ERROR_INVALID_DIMENSION,
303 SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
304 SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
305 SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
306 SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
307 SORT_ERROR_INVALID_ARGUMENT,
308 SORT_ERROR_INTERNAL,
309];
310
311const SORT_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 2] = [
312 BuiltinIntegerInputCapability {
313 name: "A",
314 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
315 availability: BuiltinIntegerInputAvailability::Documented,
316 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
317 notes: "The documented sortable input domain includes all eight real integer classes.",
318 },
319 BuiltinIntegerInputCapability {
320 name: "dim",
321 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
322 availability: BuiltinIntegerInputAvailability::Documented,
323 scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
324 notes: "The documented positive integer-scalar dimension accepts every integer class and integer-valued scalar double.",
325 },
326];
327
328const SORT_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
329 [BuiltinIntegerCapabilityDescriptor {
330 form: "[B, I] = sort(integer_A, integer_dim, direction, options)",
331 inputs: &SORT_INTEGER_INPUTS,
332 computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
333 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
334 overflow: BuiltinIntegerOverflowRule::NotApplicable,
335 backend: BuiltinIntegerBackendRule::HostAndGpu,
336 overload: BuiltinIntegerOverloadKind::Multiple,
337 notes: "B preserves A's exact integer class and stable equal-value order; optional I is one-based double. Resident integer input uses exact typed gather fallback and restores both outputs to the owning provider.",
338 }];
339
340pub const SORT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
341 signatures: &SORT_SIGNATURES,
342 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
343 completion_policy: BuiltinCompletionPolicy::Public,
344 errors: &SORT_ERRORS,
345};
346
347fn sort_error(
348 error: &'static BuiltinErrorDescriptor,
349 message: impl Into<String>,
350) -> crate::RuntimeError {
351 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
352 if let Some(identifier) = error.identifier {
353 builder = builder.with_identifier(identifier);
354 }
355 builder.build()
356}
357
358fn sort_internal(message: impl Into<String>) -> crate::RuntimeError {
359 sort_error(&SORT_ERROR_INTERNAL, message)
360}
361
362fn sort_invalid_argument(message: impl Into<String>) -> crate::RuntimeError {
363 sort_error(&SORT_ERROR_INVALID_ARGUMENT, message)
364}
365
366#[runtime_builtin(
367 name = "sort",
368 category = "array/sorting_sets",
369 summary = "Sort array elements along a dimension with optional index outputs.",
370 keywords = "sort,ascending,descending,indices,comparisonmethod,gpu",
371 accel = "sink",
372 sink = true,
373 type_resolver(tensor_output_type),
374 descriptor(crate::builtins::array::sorting_sets::sort::SORT_DESCRIPTOR),
375 integer_capabilities(SORT_INTEGER_CAPABILITIES),
376 builtin_path = "crate::builtins::array::sorting_sets::sort"
377)]
378async fn sort_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
379 let eval = evaluate(value, &rest).await?;
380 if let Some(out_count) = crate::output_count::current_output_count() {
381 if out_count == 0 {
382 return Ok(Value::OutputList(Vec::new()));
383 }
384 let (sorted, indices) = eval.into_values();
385 if out_count == 1 {
386 return Ok(Value::OutputList(vec![sorted]));
387 }
388 if out_count == 2 {
389 return Ok(Value::OutputList(vec![sorted, indices]));
390 }
391 return Err(sort_invalid_argument(
392 "sort: too many output arguments; maximum is 2",
393 ));
394 }
395 Ok(eval.into_sorted_value())
396}
397
398pub async fn evaluate(value: Value, rest: &[Value]) -> crate::BuiltinResult<SortEvaluation> {
400 let args = SortArgs::parse(rest)?;
401 match value {
402 Value::GpuTensor(handle) => sort_gpu(handle, &args).await,
403 other => sort_host(other, &args),
404 }
405}
406
407async fn sort_gpu(
408 handle: GpuTensorHandle,
409 args: &SortArgs,
410) -> crate::BuiltinResult<SortEvaluation> {
411 let shape = handle.shape.clone();
412 let provider = runmat_accelerate_api::provider_for_handle(&handle)
413 .or_else(runmat_accelerate_api::provider);
414 let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
415 if dim == 0 {
416 return Err(sort_error(
417 &SORT_ERROR_INVALID_DIMENSION,
418 "sort: dimension must be >= 1",
419 ));
420 }
421 let dim_len = dimension_length(&shape, dim);
422 let plain_real = runmat_accelerate_api::handle_integer_type(&handle).is_none()
423 && !runmat_accelerate_api::handle_is_logical(&handle)
424 && runmat_accelerate_api::handle_storage(&handle)
425 == runmat_accelerate_api::GpuTensorStorage::Real;
426 if dim_len > 1 && plain_real && matches!(args.missing, MissingPlacement::Auto) {
427 if let Some(provider) = provider {
428 let order = args.direction.to_provider();
429 let comparison = args.comparison.to_provider();
430 let zero_based = dim - 1;
431 if let Ok(result) = provider
432 .sort_dim(&handle, zero_based, order, comparison)
433 .await
434 {
435 let sorted_tensor = Tensor::new(result.values.data, result.values.shape)
436 .map_err(|e| sort_internal(format!("sort: {e}")))?;
437 let sorted_value = tensor::tensor_into_value(sorted_tensor);
438 let indices_tensor = Tensor::new(result.indices.data, result.indices.shape)
439 .map_err(|e| sort_internal(format!("sort: {e}")))?;
440 return upload_sort_evaluation(
441 provider,
442 SortEvaluation {
443 sorted: sorted_value,
444 indices: tensor::tensor_into_value(indices_tensor),
445 },
446 );
447 }
448 }
449 }
450 let host = gpu_helpers::gather_value_async(&Value::GpuTensor(handle)).await?;
451 let evaluation = sort_host(host, args)?;
452 match provider {
453 Some(provider) => upload_sort_evaluation(provider, evaluation),
454 None => Ok(evaluation),
455 }
456}
457
458fn sort_host(value: Value, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
459 match value {
460 Value::ComplexTensor(ct) => sort_complex_tensor(ct, args),
461 Value::Complex(re, im) => {
462 let tensor = ComplexTensor::new(vec![(re, im)], vec![1, 1])
463 .map_err(|e| sort_internal(format!("sort: {e}")))?;
464 sort_complex_tensor(tensor, args)
465 }
466 Value::Int(value) => {
467 let tensor = Tensor::new_integer(IntegerStorage::from_scalar(value), vec![1, 1])
468 .map_err(|e| sort_internal(format!("sort: {e}")))?;
469 sort_real_tensor(tensor, args)
470 }
471 Value::LogicalArray(logical) => sort_logical(logical, args),
472 Value::Bool(value) => {
473 let logical = LogicalArray::new(vec![u8::from(value)], vec![1, 1])
474 .map_err(|e| sort_internal(format!("sort: {e}")))?;
475 sort_logical(logical, args)
476 }
477 other => {
478 let tensor =
479 tensor::value_into_tensor_for("sort", other).map_err(sort_invalid_argument)?;
480 sort_real_tensor(tensor, args)
481 }
482 }
483}
484
485fn sort_logical(logical: LogicalArray, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
486 let shape = logical.shape;
487 let evaluation = sort_real_tensor(
488 Tensor::new_integer(IntegerStorage::U8(logical.data), shape)
489 .map_err(|e| sort_internal(format!("sort: {e}")))?,
490 args,
491 )?;
492 let sorted = match evaluation.sorted {
493 Value::Int(IntValue::U8(value)) => Value::Bool(value != 0),
494 Value::Tensor(tensor) => {
495 let shape = tensor.shape.clone();
496 let storage = tensor
497 .into_numeric_storage()
498 .map_err(|e| sort_internal(format!("sort: {e}")))?
499 .into_integer_storage()
500 .map_err(|_| sort_internal("sort: logical output lost integer storage"))?;
501 let IntegerStorage::U8(values) = storage else {
502 return Err(sort_internal("sort: logical output changed storage class"));
503 };
504 if values.len() == 1 {
505 Value::Bool(values[0] != 0)
506 } else {
507 Value::LogicalArray(
508 LogicalArray::new(values, shape)
509 .map_err(|e| sort_internal(format!("sort: {e}")))?,
510 )
511 }
512 }
513 other => {
514 return Err(sort_internal(format!(
515 "sort: unexpected logical output {other:?}"
516 )))
517 }
518 };
519 Ok(SortEvaluation {
520 sorted,
521 indices: evaluation.indices,
522 })
523}
524
525fn upload_sort_evaluation(
526 provider: &'static dyn runmat_accelerate_api::AccelProvider,
527 evaluation: SortEvaluation,
528) -> crate::BuiltinResult<SortEvaluation> {
529 Ok(SortEvaluation {
530 sorted: upload_sort_value(provider, evaluation.sorted)?,
531 indices: upload_sort_value(provider, evaluation.indices)?,
532 })
533}
534
535fn upload_sort_value(
536 provider: &'static dyn runmat_accelerate_api::AccelProvider,
537 value: Value,
538) -> crate::BuiltinResult<Value> {
539 let upload_tensor = |tensor: Tensor, logical: bool| -> crate::BuiltinResult<Value> {
540 let handle = gpu_helpers::upload_tensor(provider, &tensor)
541 .map_err(|error| sort_internal(format!("sort: GPU upload failed: {error}")))?;
542 Ok(if logical {
543 gpu_helpers::logical_gpu_value(handle)
544 } else {
545 gpu_helpers::resident_gpu_value(handle)
546 })
547 };
548 match value {
549 Value::Tensor(tensor) => upload_tensor(tensor, false),
550 Value::Num(number) => upload_tensor(
551 Tensor::new(vec![number], vec![1, 1])
552 .map_err(|error| sort_internal(format!("sort: {error}")))?,
553 false,
554 ),
555 Value::Int(integer) => upload_tensor(
556 Tensor::new_integer(IntegerStorage::from_scalar(integer), vec![1, 1])
557 .map_err(|error| sort_internal(format!("sort: {error}")))?,
558 false,
559 ),
560 Value::Bool(logical) => upload_tensor(
561 Tensor::new(vec![if logical { 1.0 } else { 0.0 }], vec![1, 1])
562 .map_err(|error| sort_internal(format!("sort: {error}")))?,
563 true,
564 ),
565 Value::LogicalArray(logical) => {
566 let tensor = tensor::logical_to_tensor(&logical).map_err(sort_internal)?;
567 upload_tensor(tensor, true)
568 }
569 Value::Complex(real, imag) => {
570 let tensor = ComplexTensor::new(vec![(real, imag)], vec![1, 1])
571 .map_err(|error| sort_internal(format!("sort: {error}")))?;
572 let handle = gpu_helpers::upload_complex_tensor(provider, &tensor)
573 .map_err(|error| sort_internal(format!("sort: {}", error.message())))?;
574 Ok(gpu_helpers::complex_gpu_value(handle))
575 }
576 Value::ComplexTensor(tensor) => {
577 let handle = gpu_helpers::upload_complex_tensor(provider, &tensor)
578 .map_err(|error| sort_internal(format!("sort: {}", error.message())))?;
579 Ok(gpu_helpers::complex_gpu_value(handle))
580 }
581 other => Err(sort_internal(format!(
582 "sort: cannot upload unexpected output {other:?}"
583 ))),
584 }
585}
586
587fn sort_real_tensor(tensor: Tensor, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
588 let shape = tensor.shape.clone();
589 let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
590 if dim == 0 {
591 return Err(sort_error(
592 &SORT_ERROR_INVALID_DIMENSION,
593 "sort: dimension must be >= 1",
594 ));
595 }
596
597 let storage = tensor
598 .into_numeric_storage()
599 .map_err(|e| sort_internal(format!("sort: {e}")))?;
600 match storage {
601 NumericStorage::F64(values) => {
602 sort_floating_tensor(values, shape, dim, args, NumericStorage::F64)
603 }
604 NumericStorage::F32(values) => {
605 sort_floating_tensor(values, shape, dim, args, NumericStorage::F32)
606 }
607 storage => sort_integer_tensor(
608 storage
609 .into_integer_storage()
610 .expect("non-floating numeric storage is integer"),
611 shape,
612 dim,
613 args,
614 ),
615 }
616}
617
618fn sort_floating_tensor<T>(
619 values: Vec<T>,
620 shape: Vec<usize>,
621 dim: usize,
622 args: &SortArgs,
623 wrap: fn(Vec<T>) -> NumericStorage,
624) -> crate::BuiltinResult<SortEvaluation>
625where
626 T: SetFloat,
627{
628 let dim_len = dimension_length(&shape, dim);
629 if values.is_empty() || dim_len <= 1 {
630 let indices = vec![1.0; values.len()];
631 let index_tensor =
632 Tensor::new(indices, shape.clone()).map_err(|e| sort_internal(format!("sort: {e}")))?;
633 let sorted_tensor = Tensor::from_numeric_storage(wrap(values), shape)
634 .map_err(|e| sort_internal(format!("sort: {e}")))?;
635 return Ok(SortEvaluation {
636 sorted: tensor::tensor_into_value(sorted_tensor),
637 indices: tensor::tensor_into_value(index_tensor),
638 });
639 }
640
641 let stride_before = stride_before(&shape, dim);
642 let stride_after = stride_after(&shape, dim);
643 let mut sorted = values;
644 let mut indices = vec![0.0f64; sorted.len()];
645 let mut buffer: Vec<(usize, T)> = Vec::with_capacity(dim_len);
646
647 for after in 0..stride_after {
648 for before in 0..stride_before {
649 buffer.clear();
650 for k in 0..dim_len {
651 let idx = before + k * stride_before + after * stride_before * dim_len;
652 let value = sorted[idx];
653 buffer.push((k, value));
654 }
655 buffer.sort_by(|a, b| compare_real_values(a.1, b.1, args));
656 for (pos, (original_index, value)) in buffer.iter().enumerate() {
657 let target = before + pos * stride_before + after * stride_before * dim_len;
658 sorted[target] = *value;
659 indices[target] = (*original_index + 1) as f64;
660 }
661 }
662 }
663
664 let sorted_tensor = Tensor::from_numeric_storage(wrap(sorted), shape.clone())
665 .map_err(|e| sort_internal(format!("sort: {e}")))?;
666 let index_tensor =
667 Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
668
669 Ok(SortEvaluation {
670 sorted: tensor::tensor_into_value(sorted_tensor),
671 indices: tensor::tensor_into_value(index_tensor),
672 })
673}
674
675fn sort_integer_tensor(
676 storage: IntegerStorage,
677 shape: Vec<usize>,
678 dim: usize,
679 args: &SortArgs,
680) -> crate::BuiltinResult<SortEvaluation> {
681 let stride_before = stride_before(&shape, dim);
682 let stride_after = stride_after(&shape, dim);
683 let dim_len = dimension_length(&shape, dim);
684 let mut sorted = storage.exact_values();
685 let mut indices = vec![0.0f64; sorted.len()];
686 let mut buffer: Vec<(usize, IntValue)> = Vec::with_capacity(dim_len);
687
688 for after in 0..stride_after {
689 for before in 0..stride_before {
690 buffer.clear();
691 for k in 0..dim_len {
692 let idx = before + k * stride_before + after * stride_before * dim_len;
693 buffer.push((k, sorted[idx].clone()));
694 }
695 buffer.sort_by(|a, b| {
696 integer_order::compare(
697 &a.1,
698 &b.1,
699 matches!(args.direction, SortDirection::Descend),
700 matches!(args.comparison, ComparisonMethod::Abs),
701 )
702 });
703 for (pos, (original_index, value)) in buffer.iter().enumerate() {
704 let target = before + pos * stride_before + after * stride_before * dim_len;
705 sorted[target] = value.clone();
706 indices[target] = (*original_index + 1) as f64;
707 }
708 }
709 }
710
711 let sorted_storage = storage
712 .from_exact_values_like(sorted)
713 .map_err(|e| sort_internal(format!("sort: {e}")))?;
714 let sorted_tensor = Tensor::new_integer(sorted_storage, shape.clone())
715 .map_err(|e| sort_internal(format!("sort: {e}")))?;
716 let index_tensor =
717 Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
718 Ok(SortEvaluation {
719 sorted: Value::Tensor(sorted_tensor),
720 indices: tensor::tensor_into_value(index_tensor),
721 })
722}
723
724fn sort_complex_tensor(
725 tensor: ComplexTensor,
726 args: &SortArgs,
727) -> crate::BuiltinResult<SortEvaluation> {
728 let shape = tensor.shape.clone();
729 let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
730 if dim == 0 {
731 return Err(sort_error(
732 &SORT_ERROR_INVALID_DIMENSION,
733 "sort: dimension must be >= 1",
734 ));
735 }
736
737 let dim_len = dimension_length(&shape, dim);
738 let storage = tensor.into_complex_storage();
739 if storage.is_empty() || dim_len <= 1 {
740 let indices = vec![1.0; storage.len()];
741 let index_tensor =
742 Tensor::new(indices, shape.clone()).map_err(|e| sort_internal(format!("sort: {e}")))?;
743 let sorted_tensor = ComplexTensor::from_complex_storage(storage, shape)
744 .map_err(|e| sort_internal(format!("sort: {e}")))?;
745 return Ok(SortEvaluation {
746 sorted: complex_tensor_into_value(sorted_tensor),
747 indices: tensor::tensor_into_value(index_tensor),
748 });
749 }
750
751 let (source_indices, indices) = match &storage {
752 ComplexStorage::F64(values) => complex_sort_permutation(values, &shape, dim, args),
753 ComplexStorage::F32(values) => complex_sort_permutation(values, &shape, dim, args),
754 ComplexStorage::Integer(values) => {
755 complex_integer_sort_permutation(values, &shape, dim, args)
756 }
757 };
758 let sorted_storage = storage
759 .gather(&source_indices)
760 .map_err(|e| sort_internal(format!("sort: {e}")))?;
761 let sorted_tensor = ComplexTensor::from_complex_storage(sorted_storage, shape.clone())
762 .map_err(|e| sort_internal(format!("sort: {e}")))?;
763 let index_tensor =
764 Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
765
766 Ok(SortEvaluation {
767 sorted: complex_tensor_into_value(sorted_tensor),
768 indices: tensor::tensor_into_value(index_tensor),
769 })
770}
771
772fn complex_integer_sort_permutation(
773 storage: &runmat_value::IntegerComplexStorage,
774 shape: &[usize],
775 dim: usize,
776 args: &SortArgs,
777) -> (Vec<usize>, Vec<f64>) {
778 let stride_before = stride_before(shape, dim);
779 let stride_after = stride_after(shape, dim);
780 let dim_len = dimension_length(shape, dim);
781 let mut source_indices = vec![0usize; storage.len()];
782 let mut indices = vec![0.0f64; storage.len()];
783 let mut buffer: Vec<(usize, usize)> = Vec::with_capacity(dim_len);
784
785 for after in 0..stride_after {
786 for before in 0..stride_before {
787 buffer.clear();
788 for k in 0..dim_len {
789 let source = before + k * stride_before + after * stride_before * dim_len;
790 buffer.push((k, source));
791 }
792 buffer.sort_by(|a, b| {
793 let a_real = storage.real.value_at(a.1).expect("validated complex index");
794 let a_imag = storage.imag.value_at(a.1).expect("validated complex index");
795 let b_real = storage.real.value_at(b.1).expect("validated complex index");
796 let b_imag = storage.imag.value_at(b.1).expect("validated complex index");
797 integer_order::compare_complex(
798 (&a_real, &a_imag),
799 (&b_real, &b_imag),
800 matches!(args.direction, SortDirection::Descend),
801 matches!(args.comparison, ComparisonMethod::Real),
802 )
803 });
804 for (position, (original_index, source)) in buffer.iter().enumerate() {
805 let target = before + position * stride_before + after * stride_before * dim_len;
806 source_indices[target] = *source;
807 indices[target] = (*original_index + 1) as f64;
808 }
809 }
810 }
811
812 (source_indices, indices)
813}
814
815fn complex_sort_permutation<T: SetFloat>(
816 values: &[(T, T)],
817 shape: &[usize],
818 dim: usize,
819 args: &SortArgs,
820) -> (Vec<usize>, Vec<f64>) {
821 let stride_before = stride_before(shape, dim);
822 let stride_after = stride_after(shape, dim);
823 let dim_len = dimension_length(shape, dim);
824 let mut source_indices = vec![0usize; values.len()];
825 let mut indices = vec![0.0f64; values.len()];
826 let mut buffer: Vec<(usize, usize, (T, T))> = Vec::with_capacity(dim_len);
827
828 for after in 0..stride_after {
829 for before in 0..stride_before {
830 buffer.clear();
831 for k in 0..dim_len {
832 let source = before + k * stride_before + after * stride_before * dim_len;
833 buffer.push((k, source, values[source]));
834 }
835 buffer.sort_by(|a, b| compare_complex_values(a.2, b.2, args));
836 for (position, (original_index, source, _)) in buffer.iter().enumerate() {
837 let target = before + position * stride_before + after * stride_before * dim_len;
838 source_indices[target] = *source;
839 indices[target] = (*original_index + 1) as f64;
840 }
841 }
842 }
843
844 (source_indices, indices)
845}
846
847fn complex_tensor_into_value(tensor: ComplexTensor) -> Value {
848 if let Some([value]) = tensor.as_f64_slice() {
849 Value::Complex(value.0, value.1)
850 } else {
851 Value::ComplexTensor(tensor)
852 }
853}
854
855fn compare_real_values<T: SetFloat>(a: T, b: T, args: &SortArgs) -> Ordering {
856 match (a.is_nan(), b.is_nan()) {
857 (true, true) => Ordering::Equal,
858 (true, false) => match args.missing.resolve(args.direction) {
859 MissingPlacementResolved::First => Ordering::Less,
860 MissingPlacementResolved::Last => Ordering::Greater,
861 },
862 (false, true) => match args.missing.resolve(args.direction) {
863 MissingPlacementResolved::First => Ordering::Greater,
864 MissingPlacementResolved::Last => Ordering::Less,
865 },
866 (false, false) => compare_real_finite(a, b, args),
867 }
868}
869
870fn compare_real_finite<T: SetFloat>(a: T, b: T, args: &SortArgs) -> Ordering {
871 let primary = match args.comparison {
872 ComparisonMethod::Abs => {
873 let abs_cmp = a.abs().compare(b.abs());
874 if abs_cmp != Ordering::Equal {
875 return match args.direction {
876 SortDirection::Ascend => abs_cmp,
877 SortDirection::Descend => abs_cmp.reverse(),
878 };
879 }
880 Ordering::Equal
881 }
882 ComparisonMethod::Auto | ComparisonMethod::Real => Ordering::Equal,
883 };
884 if primary != Ordering::Equal {
885 return primary;
886 }
887 let ordering = if matches!(args.comparison, ComparisonMethod::Abs) {
888 b.compare(a)
889 } else {
890 a.compare(b)
891 };
892 match args.direction {
893 SortDirection::Ascend => ordering,
894 SortDirection::Descend => ordering.reverse(),
895 }
896}
897
898fn compare_complex_values<T: SetFloat>(a: (T, T), b: (T, T), args: &SortArgs) -> Ordering {
899 match (complex_is_nan(a), complex_is_nan(b)) {
900 (true, true) => Ordering::Equal,
901 (true, false) => match args.missing.resolve(args.direction) {
902 MissingPlacementResolved::First => Ordering::Less,
903 MissingPlacementResolved::Last => Ordering::Greater,
904 },
905 (false, true) => match args.missing.resolve(args.direction) {
906 MissingPlacementResolved::First => Ordering::Greater,
907 MissingPlacementResolved::Last => Ordering::Less,
908 },
909 (false, false) => compare_complex_finite(a, b, args),
910 }
911}
912
913fn compare_complex_finite<T: SetFloat>(a: (T, T), b: (T, T), args: &SortArgs) -> Ordering {
914 match args.comparison {
915 ComparisonMethod::Real => compare_complex_real_imag(a, b, args.direction),
916 ComparisonMethod::Abs | ComparisonMethod::Auto => {
917 let abs_cmp = complex_abs(a).compare(complex_abs(b));
918 if abs_cmp != Ordering::Equal {
919 return match args.direction {
920 SortDirection::Ascend => abs_cmp,
921 SortDirection::Descend => abs_cmp.reverse(),
922 };
923 }
924 compare_complex_phase(a, b, args.direction)
925 }
926 }
927}
928
929fn compare_complex_phase<T: SetFloat>(a: (T, T), b: (T, T), direction: SortDirection) -> Ordering {
930 let ordering = complex_phase(a).compare(complex_phase(b));
931 match direction {
932 SortDirection::Ascend => ordering,
933 SortDirection::Descend => ordering.reverse(),
934 }
935}
936
937fn complex_phase<T: SetFloat>((real, imaginary): (T, T)) -> T {
938 let imaginary = if imaginary == T::default() {
939 T::default()
940 } else {
941 imaginary
942 };
943 imaginary.atan2(real)
944}
945
946fn compare_complex_real_imag<T: SetFloat>(
947 a: (T, T),
948 b: (T, T),
949 direction: SortDirection,
950) -> Ordering {
951 let real_cmp = match direction {
952 SortDirection::Ascend => a.0.compare(b.0),
953 SortDirection::Descend => b.0.compare(a.0),
954 };
955 if real_cmp != Ordering::Equal {
956 return real_cmp;
957 }
958 match direction {
959 SortDirection::Ascend => a.1.compare(b.1),
960 SortDirection::Descend => b.1.compare(a.1),
961 }
962}
963
964fn complex_is_nan<T: SetFloat>(value: (T, T)) -> bool {
965 value.0.is_nan() || value.1.is_nan()
966}
967
968fn complex_abs<T: SetFloat>(value: (T, T)) -> T {
969 value.0.hypot(value.1)
970}
971
972fn stride_before(shape: &[usize], dim: usize) -> usize {
973 if dim <= 1 {
974 return 1;
975 }
976 let mut product = 1usize;
977 for i in 0..(dim - 1) {
978 product = product.saturating_mul(*shape.get(i).unwrap_or(&1));
979 }
980 product
981}
982
983fn stride_after(shape: &[usize], dim: usize) -> usize {
984 if dim >= shape.len() {
985 return 1;
986 }
987 let mut product = 1usize;
988 for extent in shape.iter().skip(dim) {
989 product = product.saturating_mul(*extent);
990 }
991 product
992}
993
994fn dimension_length(shape: &[usize], dim: usize) -> usize {
995 shape.get(dim - 1).copied().unwrap_or(1)
996}
997
998fn default_dimension(shape: &[usize]) -> usize {
999 shape
1000 .iter()
1001 .position(|&extent| extent > 1)
1002 .map(|idx| idx + 1)
1003 .unwrap_or(1)
1004}
1005
1006#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1007enum SortDirection {
1008 #[default]
1009 Ascend,
1010 Descend,
1011}
1012
1013impl SortDirection {
1014 fn to_provider(self) -> ProviderSortOrder {
1015 match self {
1016 SortDirection::Ascend => ProviderSortOrder::Ascend,
1017 SortDirection::Descend => ProviderSortOrder::Descend,
1018 }
1019 }
1020}
1021
1022#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1023enum ComparisonMethod {
1024 #[default]
1025 Auto,
1026 Real,
1027 Abs,
1028}
1029
1030#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1031enum MissingPlacement {
1032 #[default]
1033 Auto,
1034 First,
1035 Last,
1036}
1037
1038#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1039enum MissingPlacementResolved {
1040 First,
1041 Last,
1042}
1043
1044impl MissingPlacement {
1045 fn resolve(self, direction: SortDirection) -> MissingPlacementResolved {
1046 match self {
1047 Self::First => MissingPlacementResolved::First,
1048 Self::Last => MissingPlacementResolved::Last,
1049 Self::Auto => match direction {
1050 SortDirection::Ascend => MissingPlacementResolved::Last,
1051 SortDirection::Descend => MissingPlacementResolved::First,
1052 },
1053 }
1054 }
1055}
1056
1057impl ComparisonMethod {
1058 fn to_provider(self) -> ProviderSortComparison {
1059 match self {
1060 ComparisonMethod::Auto => ProviderSortComparison::Auto,
1061 ComparisonMethod::Real => ProviderSortComparison::Real,
1062 ComparisonMethod::Abs => ProviderSortComparison::Abs,
1063 }
1064 }
1065}
1066
1067#[derive(Debug, Clone, Default)]
1068struct SortArgs {
1069 dimension: Option<usize>,
1070 direction: SortDirection,
1071 comparison: ComparisonMethod,
1072 missing: MissingPlacement,
1073}
1074
1075impl SortArgs {
1076 fn parse(rest: &[Value]) -> crate::BuiltinResult<Self> {
1077 let mut args = SortArgs::default();
1078 let tokens = tokens_from_values(rest);
1079 let mut i = 0usize;
1080 while i < rest.len() {
1081 if args.dimension.is_none() {
1082 if is_dimension_placeholder(&rest[i]) {
1083 i += 1;
1084 continue;
1085 }
1086 match tensor::parse_dimension(&rest[i], "sort") {
1087 Ok(dim) => {
1088 args.dimension = Some(dim);
1089 i += 1;
1090 continue;
1091 }
1092 Err(err) => {
1093 if matches!(rest[i], Value::Int(_) | Value::Num(_)) {
1094 return Err(sort_error(&SORT_ERROR_INVALID_DIMENSION, err));
1095 }
1096 }
1097 }
1098 }
1099 if let Some(ArgToken::String(text)) = tokens.get(i) {
1100 match text.as_str() {
1101 "ascend" | "ascending" => {
1102 args.direction = SortDirection::Ascend;
1103 i += 1;
1104 continue;
1105 }
1106 "descend" | "descending" => {
1107 args.direction = SortDirection::Descend;
1108 i += 1;
1109 continue;
1110 }
1111 "comparisonmethod" => {
1112 i += 1;
1113 if i >= rest.len() {
1114 return Err(sort_invalid_argument(
1115 "sort: expected a value for 'ComparisonMethod'",
1116 ));
1117 }
1118 let value = match tokens.get(i) {
1119 Some(ArgToken::String(value)) => value.as_str(),
1120 _ => {
1121 return Err(sort_error(
1122 &SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
1123 SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.message,
1124 ))
1125 }
1126 };
1127 args.comparison = match value {
1128 "auto" => ComparisonMethod::Auto,
1129 "real" => ComparisonMethod::Real,
1130 "abs" | "magnitude" => ComparisonMethod::Abs,
1131 other => {
1132 return Err(sort_error(
1133 &SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
1134 format!("sort: unsupported ComparisonMethod '{other}'"),
1135 )
1136 .into())
1137 }
1138 };
1139 i += 1;
1140 continue;
1141 }
1142 "missingplacement" => {
1143 i += 1;
1144 if i >= rest.len() {
1145 return Err(sort_invalid_argument(
1146 "sort: expected a value for 'MissingPlacement'",
1147 ));
1148 }
1149 let value = match tokens.get(i) {
1150 Some(ArgToken::String(value)) => value.as_str(),
1151 _ => {
1152 return Err(sort_error(
1153 &SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
1154 SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING.message,
1155 ))
1156 }
1157 };
1158 args.missing = match value {
1159 "auto" => MissingPlacement::Auto,
1160 "first" => MissingPlacement::First,
1161 "last" => MissingPlacement::Last,
1162 other => {
1163 return Err(sort_error(
1164 &SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
1165 format!("sort: unsupported MissingPlacement '{other}'"),
1166 ))
1167 }
1168 };
1169 i += 1;
1170 continue;
1171 }
1172 _ => {}
1173 }
1174 }
1175 if let Some(keyword) = tensor::value_to_string(&rest[i]) {
1176 let lowered = keyword.trim().to_ascii_lowercase();
1177 match lowered.as_str() {
1178 "ascend" | "ascending" => {
1179 args.direction = SortDirection::Ascend;
1180 i += 1;
1181 continue;
1182 }
1183 "descend" | "descending" => {
1184 args.direction = SortDirection::Descend;
1185 i += 1;
1186 continue;
1187 }
1188 "comparisonmethod" => {
1189 i += 1;
1190 if i >= rest.len() {
1191 return Err(sort_invalid_argument(
1192 "sort: expected a value for 'ComparisonMethod'",
1193 ));
1194 }
1195 let raw = &rest[i];
1196 let value = match raw {
1197 Value::String(s) => s.clone(),
1198 Value::StringArray(sa) if sa.data.len() == 1 => sa.data[0].clone(),
1199 Value::CharArray(ca) if ca.rows == 1 => {
1200 ca.data.iter().copied().collect()
1201 }
1202 _ => {
1203 return Err(sort_error(
1204 &SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
1205 SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.message,
1206 ))
1207 }
1208 };
1209 let lowered_value = value.trim().to_ascii_lowercase();
1210 args.comparison = match lowered_value.as_str() {
1211 "auto" => ComparisonMethod::Auto,
1212 "real" => ComparisonMethod::Real,
1213 "abs" | "magnitude" => ComparisonMethod::Abs,
1214 other => {
1215 return Err(sort_error(
1216 &SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
1217 format!("sort: unsupported ComparisonMethod '{other}'"),
1218 )
1219 .into())
1220 }
1221 };
1222 i += 1;
1223 continue;
1224 }
1225 "missingplacement" => {
1226 i += 1;
1227 if i >= rest.len() {
1228 return Err(sort_invalid_argument(
1229 "sort: expected a value for 'MissingPlacement'",
1230 ));
1231 }
1232 let value = tensor::value_to_string(&rest[i]).ok_or_else(|| {
1233 sort_error(
1234 &SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
1235 SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING.message,
1236 )
1237 })?;
1238 args.missing = match value.trim().to_ascii_lowercase().as_str() {
1239 "auto" => MissingPlacement::Auto,
1240 "first" => MissingPlacement::First,
1241 "last" => MissingPlacement::Last,
1242 other => {
1243 return Err(sort_error(
1244 &SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
1245 format!("sort: unsupported MissingPlacement '{other}'"),
1246 ))
1247 }
1248 };
1249 i += 1;
1250 continue;
1251 }
1252 _ => {}
1253 }
1254 }
1255 return Err(sort_invalid_argument(format!(
1256 "sort: unrecognised argument {:?}",
1257 rest[i]
1258 )));
1259 }
1260 Ok(args)
1261 }
1262}
1263
1264fn is_dimension_placeholder(value: &Value) -> bool {
1265 match value {
1266 Value::Tensor(t) => tensor::tensor_element_len(t) == 0,
1267 Value::LogicalArray(logical) => logical.data.is_empty(),
1268 _ => false,
1269 }
1270}
1271
1272pub struct SortEvaluation {
1273 sorted: Value,
1274 indices: Value,
1275}
1276
1277impl SortEvaluation {
1278 pub fn into_sorted_value(self) -> Value {
1279 self.sorted
1280 }
1281
1282 pub fn into_values(self) -> (Value, Value) {
1283 (self.sorted, self.indices)
1284 }
1285
1286 pub fn indices_value(&self) -> Value {
1287 self.indices.clone()
1288 }
1289}
1290
1291#[cfg(test)]
1292pub(crate) mod tests {
1293 use super::*;
1294 use crate::builtins::common::test_support;
1295 use futures::executor::block_on;
1296 use runmat_builtins::{ResolveContext, Type};
1297 use runmat_value::{
1298 ComplexTensor, IntValue, IntegerComplexStorage, IntegerStorage, NumericStorage, Tensor,
1299 Value,
1300 };
1301
1302 fn sort_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1303 block_on(super::sort_builtin(value, rest))
1304 }
1305
1306 fn evaluate(value: Value, rest: &[Value]) -> crate::BuiltinResult<SortEvaluation> {
1307 block_on(super::evaluate(value, rest))
1308 }
1309
1310 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1311 #[test]
1312 fn sort_vector_default() {
1313 let tensor = Tensor::new(vec![3.0, 1.0, 2.0], vec![3, 1]).unwrap();
1314 let result = sort_builtin(Value::Tensor(tensor), Vec::new()).expect("sort");
1315 match result {
1316 Value::Tensor(t) => {
1317 assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0]);
1318 assert_eq!(t.shape, vec![3, 1]);
1319 }
1320 other => panic!("expected tensor result, got {other:?}"),
1321 }
1322 }
1323
1324 #[test]
1325 fn sort_preserves_native_single_storage_and_indices() {
1326 let tensor = Tensor::from_f32(vec![3.0, 1.0, 2.0], vec![3, 1]).expect("input");
1327 let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1328 .expect("sort")
1329 .into_values();
1330 let Value::Tensor(sorted) = sorted else {
1331 panic!("expected sorted tensor");
1332 };
1333 assert_eq!(
1334 sorted.into_numeric_storage().expect("storage"),
1335 NumericStorage::F32(vec![1.0, 2.0, 3.0])
1336 );
1337 let Value::Tensor(indices) = indices else {
1338 panic!("expected index tensor");
1339 };
1340 assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 1.0]);
1341 }
1342
1343 #[test]
1344 fn sort_type_resolver_tensor() {
1345 assert_eq!(
1346 tensor_output_type(&[Type::tensor()], &ResolveContext::new(Vec::new())),
1347 Type::tensor()
1348 );
1349 }
1350
1351 #[test]
1352 fn sort_preserves_all_exact_real_integer_classes_and_indices() {
1353 let cases = [
1354 (
1355 IntegerStorage::I8(vec![i8::MAX, i8::MIN, 0, 7]),
1356 IntegerStorage::I8(vec![i8::MIN, 0, 7, i8::MAX]),
1357 ),
1358 (
1359 IntegerStorage::I16(vec![i16::MAX, i16::MIN, 0, 7]),
1360 IntegerStorage::I16(vec![i16::MIN, 0, 7, i16::MAX]),
1361 ),
1362 (
1363 IntegerStorage::I32(vec![i32::MAX, i32::MIN, 0, 7]),
1364 IntegerStorage::I32(vec![i32::MIN, 0, 7, i32::MAX]),
1365 ),
1366 (
1367 IntegerStorage::I64(vec![i64::MAX, i64::MIN, 0, 7]),
1368 IntegerStorage::I64(vec![i64::MIN, 0, 7, i64::MAX]),
1369 ),
1370 (
1371 IntegerStorage::U8(vec![u8::MAX, 0, 7, 9]),
1372 IntegerStorage::U8(vec![0, 7, 9, u8::MAX]),
1373 ),
1374 (
1375 IntegerStorage::U16(vec![u16::MAX, 0, 7, 700]),
1376 IntegerStorage::U16(vec![0, 7, 700, u16::MAX]),
1377 ),
1378 (
1379 IntegerStorage::U32(vec![u32::MAX, 0, 7, 9_007_199]),
1380 IntegerStorage::U32(vec![0, 7, 9_007_199, u32::MAX]),
1381 ),
1382 (
1383 IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1384 IntegerStorage::U64(vec![0, 7, 9_007_199_254_740_993, u64::MAX]),
1385 ),
1386 ];
1387
1388 for (input, expected) in cases {
1389 let tensor = Tensor::new_integer(input, vec![4, 1]).expect("input");
1390 let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1391 .expect("sort")
1392 .into_values();
1393 let Value::Tensor(sorted) = sorted else {
1394 panic!("expected exact integer sorted values");
1395 };
1396 assert_eq!(sorted.integer_storage(), Some(&expected));
1397 let Value::Tensor(indices) = indices else {
1398 panic!("expected index tensor");
1399 };
1400 assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 4.0, 1.0]);
1401 }
1402 }
1403
1404 #[test]
1405 fn sort_reads_exact_integer_values_without_mirror() {
1406 let tensor = Tensor::new_integer(
1407 IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993, 7]),
1408 vec![4, 1],
1409 )
1410 .expect("input");
1411
1412 let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1413 .expect("sort")
1414 .into_values();
1415 let Value::Tensor(sorted) = sorted else {
1416 panic!("expected exact integer sorted values");
1417 };
1418 assert_eq!(
1419 sorted.integer_storage(),
1420 Some(&IntegerStorage::U64(vec![
1421 0,
1422 7,
1423 9_007_199_254_740_993,
1424 u64::MAX,
1425 ]))
1426 );
1427 let Value::Tensor(indices) = indices else {
1428 panic!("expected index tensor");
1429 };
1430 assert_eq!(indices.materialize_f64(), vec![2.0, 4.0, 3.0, 1.0]);
1431 }
1432
1433 #[test]
1434 fn sort_exact_integer_honors_descend_abs_and_dimension() {
1435 let input = Tensor::new_integer(
1436 IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1437 vec![4, 1],
1438 )
1439 .expect("input");
1440 let (sorted, indices) = evaluate(Value::Tensor(input), &[Value::from("descend")])
1441 .expect("descending sort")
1442 .into_values();
1443 let Value::Tensor(sorted) = sorted else {
1444 panic!("expected exact integer sorted values");
1445 };
1446 assert_eq!(
1447 sorted.integer_storage(),
1448 Some(&IntegerStorage::U64(vec![
1449 u64::MAX,
1450 9_007_199_254_740_993,
1451 7,
1452 0,
1453 ]))
1454 );
1455 let Value::Tensor(indices) = indices else {
1456 panic!("expected index tensor");
1457 };
1458 assert_eq!(indices.materialize_f64(), vec![1.0, 4.0, 3.0, 2.0]);
1459
1460 let input = Tensor::new_integer(
1461 IntegerStorage::I64(vec![i64::MIN, i64::MAX, 2, -1]),
1462 vec![4, 1],
1463 )
1464 .expect("input");
1465 let (sorted, indices) = evaluate(
1466 Value::Tensor(input),
1467 &[Value::from("ComparisonMethod"), Value::from("abs")],
1468 )
1469 .expect("absolute sort")
1470 .into_values();
1471 let Value::Tensor(sorted) = sorted else {
1472 panic!("expected exact integer sorted values");
1473 };
1474 assert_eq!(
1475 sorted.integer_storage(),
1476 Some(&IntegerStorage::I64(vec![-1, 2, i64::MAX, i64::MIN]))
1477 );
1478 let Value::Tensor(indices) = indices else {
1479 panic!("expected index tensor");
1480 };
1481 assert_eq!(indices.materialize_f64(), vec![4.0, 3.0, 2.0, 1.0]);
1482
1483 let input = Tensor::new_integer(
1484 IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1485 vec![2, 2],
1486 )
1487 .expect("matrix input");
1488 let (sorted, indices) = evaluate(Value::Tensor(input), &[Value::Int(IntValue::I32(2))])
1489 .expect("dimension sort")
1490 .into_values();
1491 let Value::Tensor(sorted) = sorted else {
1492 panic!("expected exact integer matrix values");
1493 };
1494 assert_eq!(
1495 sorted.integer_storage(),
1496 Some(&IntegerStorage::U64(vec![
1497 7,
1498 0,
1499 u64::MAX,
1500 9_007_199_254_740_993,
1501 ]))
1502 );
1503 let Value::Tensor(indices) = indices else {
1504 panic!("expected index tensor");
1505 };
1506 assert_eq!(indices.materialize_f64(), vec![2.0, 1.0, 1.0, 2.0]);
1507 }
1508
1509 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1510 #[test]
1511 fn sort_descend_direction() {
1512 let tensor = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![4, 1]).unwrap();
1513 let result =
1514 sort_builtin(Value::Tensor(tensor), vec![Value::from("descend")]).expect("sort");
1515 match result {
1516 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, 3.0, 2.0, 1.0]),
1517 other => panic!("expected tensor, got {other:?}"),
1518 }
1519 }
1520
1521 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1522 #[test]
1523 fn sort_matrix_default_dim1() {
1524 let tensor = Tensor::new(vec![4.0, 2.0, 1.0, 5.0, 6.0, 3.0], vec![2, 3]).unwrap();
1525 let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1526 let (sorted, indices) = eval.into_values();
1527 match sorted {
1528 Value::Tensor(t) => {
1529 assert_eq!(t.materialize_f64(), vec![2.0, 4.0, 1.0, 5.0, 3.0, 6.0]);
1530 assert_eq!(t.shape, vec![2, 3]);
1531 }
1532 other => panic!("expected tensor result, got {other:?}"),
1533 }
1534 match indices {
1535 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0, 1.0, 1.0, 2.0, 2.0, 1.0]),
1536 other => panic!("expected tensor indices, got {other:?}"),
1537 }
1538 }
1539
1540 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1541 #[test]
1542 fn sort_matrix_along_dimension_two() {
1543 let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1544 let eval =
1545 evaluate(Value::Tensor(tensor), &[Value::Int(IntValue::I32(2))]).expect("evaluate");
1546 let (sorted, indices) = eval.into_values();
1547 match sorted {
1548 Value::Tensor(t) => {
1549 assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 3.0, 4.0, 5.0]);
1550 assert_eq!(t.shape, vec![2, 3]);
1551 }
1552 other => panic!("expected tensor result, got {other:?}"),
1553 }
1554 match indices {
1555 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0]),
1556 other => panic!("expected tensor indices, got {other:?}"),
1557 }
1558 }
1559
1560 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1561 #[test]
1562 fn sort_dimension_placeholder_then_dim() {
1563 let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0], vec![2, 2]).unwrap();
1564 let placeholder = Tensor::new(Vec::new(), vec![0, 0]).unwrap();
1565 let eval = evaluate(
1566 Value::Tensor(tensor),
1567 &[
1568 Value::Tensor(placeholder),
1569 Value::Int(IntValue::I32(2)),
1570 Value::from("descend"),
1571 ],
1572 )
1573 .expect("evaluate");
1574 let (sorted, _) = eval.into_values();
1575 match sorted {
1576 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, 3.0, 1.0, 2.0]),
1577 other => panic!("expected tensor result, got {other:?}"),
1578 }
1579 }
1580
1581 #[test]
1582 fn sort_typed_integer_dimension_argument_is_not_empty_placeholder_without_mirror() {
1583 let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1584 let dim = Tensor::new_integer(IntegerStorage::U16(vec![2]), vec![1, 1]).expect("dimension");
1585
1586 let eval = evaluate(Value::Tensor(tensor), &[Value::Tensor(dim)]).expect("evaluate");
1587 let (sorted, indices) = eval.into_values();
1588 match sorted {
1589 Value::Tensor(t) => {
1590 assert_eq!(t.shape, vec![2, 3]);
1591 assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 3.0, 4.0, 5.0]);
1592 }
1593 other => panic!("expected tensor result, got {other:?}"),
1594 }
1595 match indices {
1596 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0]),
1597 other => panic!("expected tensor indices, got {other:?}"),
1598 }
1599 }
1600
1601 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1602 #[test]
1603 fn sort_descend_then_dimension() {
1604 let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1605 let eval = evaluate(
1606 Value::Tensor(tensor),
1607 &[Value::from("descend"), Value::Int(IntValue::I32(1))],
1608 )
1609 .expect("evaluate");
1610 let (sorted, _) = eval.into_values();
1611 match sorted {
1612 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 4.0, 2.0, 5.0, 2.0]),
1613 other => panic!("expected tensor result, got {other:?}"),
1614 }
1615 }
1616
1617 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1618 #[test]
1619 fn sort_returns_indices() {
1620 let tensor = Tensor::new(vec![4.0, 1.0, 9.0, 2.0], vec![4, 1]).unwrap();
1621 let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1622 let (sorted, indices) = eval.into_values();
1623 match sorted {
1624 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 4.0, 9.0]),
1625 other => panic!("expected tensor, got {other:?}"),
1626 }
1627 match indices {
1628 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0, 4.0, 1.0, 3.0]),
1629 other => panic!("expected tensor, got {other:?}"),
1630 }
1631 }
1632
1633 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1634 #[test]
1635 fn sort_with_nan_handling() {
1636 let tensor = Tensor::new(vec![f64::NAN, 4.0, 1.0, 2.0], vec![4, 1]).unwrap();
1637 let eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("evaluate");
1638 let (sorted, _) = eval.into_values();
1639 match sorted {
1640 Value::Tensor(t) => {
1641 assert!(t.materialize_f64()[3].is_nan());
1642 assert_eq!(&t.materialize_f64()[0..3], &[1.0, 2.0, 4.0]);
1643 }
1644 other => panic!("expected tensor, got {other:?}"),
1645 }
1646
1647 let eval_desc =
1648 evaluate(Value::Tensor(tensor), &[Value::from("descend")]).expect("evaluate");
1649 let (sorted_desc, _) = eval_desc.into_values();
1650 match sorted_desc {
1651 Value::Tensor(t) => {
1652 assert!(t.materialize_f64()[0].is_nan());
1653 assert_eq!(&t.materialize_f64()[1..], &[4.0, 2.0, 1.0]);
1654 }
1655 other => panic!("expected tensor, got {other:?}"),
1656 }
1657 }
1658
1659 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1660 #[test]
1661 fn sort_by_absolute_value() {
1662 let tensor = Tensor::new(vec![-8.0, -1.0, 3.0, -2.0], vec![4, 1]).unwrap();
1663 let eval = evaluate(
1664 Value::Tensor(tensor),
1665 &[Value::from("ComparisonMethod"), Value::from("abs")],
1666 )
1667 .expect("evaluate");
1668 let (sorted, _) = eval.into_values();
1669 match sorted {
1670 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![-1.0, -2.0, 3.0, -8.0]),
1671 other => panic!("expected tensor, got {other:?}"),
1672 }
1673 }
1674
1675 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1676 #[test]
1677 fn sort_by_absolute_value_descend() {
1678 let tensor = Tensor::new(vec![-1.0, 2.0, -3.0, 4.0], vec![4, 1]).unwrap();
1679 let eval = evaluate(
1680 Value::Tensor(tensor),
1681 &[
1682 Value::from("descend"),
1683 Value::from("ComparisonMethod"),
1684 Value::from("abs"),
1685 ],
1686 )
1687 .expect("evaluate");
1688 let (sorted, _) = eval.into_values();
1689 match sorted {
1690 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, -3.0, 2.0, -1.0]),
1691 other => panic!("expected tensor, got {other:?}"),
1692 }
1693 }
1694
1695 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1696 #[test]
1697 fn sort_complex_auto_abs() {
1698 let tensor =
1699 ComplexTensor::new(vec![(1.0, 2.0), (-3.0, 0.5), (0.0, -1.0)], vec![3, 1]).unwrap();
1700 let eval = evaluate(Value::ComplexTensor(tensor), &[]).expect("evaluate");
1701 let (sorted, indices) = eval.into_values();
1702 match sorted {
1703 Value::ComplexTensor(t) => {
1704 assert_eq!(
1705 t.materialize_f64(),
1706 vec![(0.0, -1.0), (1.0, 2.0), (-3.0, 0.5)]
1707 )
1708 }
1709 other => panic!("expected complex tensor, got {other:?}"),
1710 }
1711 match indices {
1712 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 2.0]),
1713 other => panic!("expected tensor indices, got {other:?}"),
1714 }
1715 }
1716
1717 #[test]
1718 fn sort_preserves_native_complex_single_storage_and_indices() {
1719 let tensor =
1720 ComplexTensor::from_f32(vec![(3.0, 4.0), (1.0, 0.0), (0.0, 2.0)], vec![3, 1]).unwrap();
1721 let (sorted, indices) = evaluate(Value::ComplexTensor(tensor), &[])
1722 .expect("sort")
1723 .into_values();
1724 let Value::ComplexTensor(sorted) = sorted else {
1725 panic!("expected complex single tensor");
1726 };
1727 assert_eq!(
1728 sorted.as_f32_slice(),
1729 Some(&[(1.0, 0.0), (0.0, 2.0), (3.0, 4.0)][..])
1730 );
1731 let Value::Tensor(indices) = indices else {
1732 panic!("expected index tensor");
1733 };
1734 assert_eq!(indices.as_f64_slice(), Some(&[2.0, 3.0, 1.0][..]));
1735 }
1736
1737 #[test]
1738 fn sort_preserves_one_element_complex_single_tensor() {
1739 let tensor = ComplexTensor::from_f32(vec![(1.25, -2.5)], vec![1, 1]).unwrap();
1740 let sorted = evaluate(Value::ComplexTensor(tensor), &[])
1741 .expect("sort")
1742 .into_sorted_value();
1743 let Value::ComplexTensor(sorted) = sorted else {
1744 panic!("expected complex single tensor");
1745 };
1746 assert_eq!(sorted.as_f32_slice(), Some(&[(1.25, -2.5)][..]));
1747 }
1748
1749 #[test]
1750 fn sort_orders_typed_complex_uint64_exactly_and_preserves_storage() {
1751 for input in [
1752 (
1753 IntegerStorage::I8(vec![3, 1]),
1754 IntegerStorage::I8(vec![4, 0]),
1755 ),
1756 (
1757 IntegerStorage::I16(vec![3, 1]),
1758 IntegerStorage::I16(vec![4, 0]),
1759 ),
1760 (
1761 IntegerStorage::I32(vec![3, 1]),
1762 IntegerStorage::I32(vec![4, 0]),
1763 ),
1764 (
1765 IntegerStorage::I64(vec![3, 1]),
1766 IntegerStorage::I64(vec![4, 0]),
1767 ),
1768 (
1769 IntegerStorage::U8(vec![3, 1]),
1770 IntegerStorage::U8(vec![4, 0]),
1771 ),
1772 (
1773 IntegerStorage::U16(vec![3, 1]),
1774 IntegerStorage::U16(vec![4, 0]),
1775 ),
1776 (
1777 IntegerStorage::U32(vec![3, 1]),
1778 IntegerStorage::U32(vec![4, 0]),
1779 ),
1780 (
1781 IntegerStorage::U64(vec![3, 1]),
1782 IntegerStorage::U64(vec![4, 0]),
1783 ),
1784 ] {
1785 let expected_real = input
1786 .0
1787 .from_exact_values_like(vec![
1788 input.0.value_at(1).unwrap(),
1789 input.0.value_at(0).unwrap(),
1790 ])
1791 .unwrap();
1792 let expected_imag = input
1793 .1
1794 .from_exact_values_like(vec![
1795 input.1.value_at(1).unwrap(),
1796 input.1.value_at(0).unwrap(),
1797 ])
1798 .unwrap();
1799 let tensor = ComplexTensor::new_integer(
1800 IntegerComplexStorage::new(input.0, input.1).unwrap(),
1801 vec![2, 1],
1802 )
1803 .unwrap();
1804 let sorted = evaluate(Value::ComplexTensor(tensor), &[])
1805 .expect("typed complex integer sort")
1806 .into_sorted_value();
1807 let Value::ComplexTensor(sorted) = sorted else {
1808 panic!("expected typed complex integer tensor");
1809 };
1810 assert_eq!(
1811 sorted.integer_storage(),
1812 Some(&IntegerComplexStorage::new(expected_real, expected_imag).unwrap())
1813 );
1814 }
1815
1816 let input = IntegerComplexStorage::new(
1817 IntegerStorage::U64(vec![u64::MAX, u64::MAX - 1, 0]),
1818 IntegerStorage::U64(vec![0, 1, u64::MAX]),
1819 )
1820 .unwrap();
1821 let tensor = ComplexTensor::new_integer(input, vec![3, 1]).unwrap();
1822 let (sorted, indices) = evaluate(Value::ComplexTensor(tensor), &[])
1823 .expect("typed complex integer sort")
1824 .into_values();
1825 let Value::ComplexTensor(sorted) = sorted else {
1826 panic!("expected typed complex integer tensor");
1827 };
1828 assert_eq!(
1829 sorted.integer_storage(),
1830 Some(
1831 &IntegerComplexStorage::new(
1832 IntegerStorage::U64(vec![u64::MAX - 1, u64::MAX, 0]),
1833 IntegerStorage::U64(vec![1, 0, u64::MAX]),
1834 )
1835 .unwrap()
1836 )
1837 );
1838 let Value::Tensor(indices) = indices else {
1839 panic!("expected double indices");
1840 };
1841 assert_eq!(indices.as_f64_slice(), Some(&[2.0, 1.0, 3.0][..]));
1842 }
1843
1844 #[test]
1845 fn sort_restores_exact_typed_complex_integer_results_to_the_input_provider() {
1846 test_support::with_test_provider(|provider| {
1847 let input = ComplexTensor::new_integer(
1848 IntegerComplexStorage::new(
1849 IntegerStorage::U64(vec![u64::MAX, u64::MAX - 1]),
1850 IntegerStorage::U64(vec![0, 1]),
1851 )
1852 .unwrap(),
1853 vec![2, 1],
1854 )
1855 .unwrap();
1856 let handle =
1857 crate::builtins::common::gpu_helpers::upload_complex_tensor(provider, &input)
1858 .expect("typed complex upload");
1859 let (sorted, indices) = evaluate(Value::GpuTensor(handle), &[])
1860 .expect("resident typed complex sort")
1861 .into_values();
1862 let sorted = block_on(crate::builtins::common::gpu_helpers::gather_value_async(
1863 &sorted,
1864 ))
1865 .expect("gather sorted values");
1866 let indices = test_support::gather(indices).expect("gather sorted indices");
1867 let Value::ComplexTensor(sorted) = sorted else {
1868 panic!("expected typed complex integer tensor");
1869 };
1870 assert_eq!(
1871 sorted.integer_storage(),
1872 Some(
1873 &IntegerComplexStorage::new(
1874 IntegerStorage::U64(vec![u64::MAX - 1, u64::MAX]),
1875 IntegerStorage::U64(vec![1, 0]),
1876 )
1877 .unwrap()
1878 )
1879 );
1880 assert_eq!(indices.as_f64_slice(), Some(&[2.0, 1.0][..]));
1881 });
1882 }
1883
1884 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1885 #[test]
1886 fn sort_complex_real_descend() {
1887 let tensor =
1888 ComplexTensor::new(vec![(1.0, 2.0), (-3.0, 0.0), (1.0, -1.0)], vec![3, 1]).unwrap();
1889 let eval = evaluate(
1890 Value::ComplexTensor(tensor),
1891 &[
1892 Value::from("descend"),
1893 Value::from("ComparisonMethod"),
1894 Value::from("real"),
1895 ],
1896 )
1897 .expect("evaluate");
1898 let (sorted, _) = eval.into_values();
1899 match sorted {
1900 Value::ComplexTensor(t) => {
1901 assert_eq!(
1902 t.materialize_f64(),
1903 vec![(1.0, 2.0), (1.0, -1.0), (-3.0, 0.0)]
1904 );
1905 }
1906 other => panic!("expected complex tensor, got {other:?}"),
1907 }
1908 }
1909
1910 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1911 #[test]
1912 fn sort_stable_with_duplicates() {
1913 let tensor = Tensor::new(vec![2.0, 2.0, 1.0, 2.0], vec![4, 1]).unwrap();
1914 let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1915 let (sorted, indices) = eval.into_values();
1916 match sorted {
1917 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 2.0]),
1918 other => panic!("expected tensor, got {other:?}"),
1919 }
1920 match indices {
1921 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 2.0, 4.0]),
1922 other => panic!("expected tensor indices, got {other:?}"),
1923 }
1924 }
1925
1926 #[test]
1927 fn sort_abs_ties_follow_phase_for_real_integer_and_complex_values() {
1928 let args = [Value::from("ComparisonMethod"), Value::from("abs")];
1929 let (real, real_indices) = evaluate(
1930 Value::Tensor(Tensor::new(vec![-3.0, 3.0], vec![2, 1]).expect("real")),
1931 &args,
1932 )
1933 .expect("real sort")
1934 .into_values();
1935 let Value::Tensor(real) = real else {
1936 panic!("expected real tensor");
1937 };
1938 assert_eq!(real.materialize_f64(), vec![3.0, -3.0]);
1939 let Value::Tensor(real_indices) = real_indices else {
1940 panic!("expected real indices");
1941 };
1942 assert_eq!(real_indices.materialize_f64(), vec![2.0, 1.0]);
1943
1944 let (integer, _) = evaluate(
1945 Value::Tensor(
1946 Tensor::new_integer(IntegerStorage::I64(vec![-3, 3]), vec![2, 1]).expect("integer"),
1947 ),
1948 &args,
1949 )
1950 .expect("integer sort")
1951 .into_values();
1952 let Value::Tensor(integer) = integer else {
1953 panic!("expected integer tensor");
1954 };
1955 assert_eq!(
1956 integer.integer_storage(),
1957 Some(&IntegerStorage::I64(vec![3, -3]))
1958 );
1959
1960 let (complex, complex_indices) = evaluate(
1961 Value::ComplexTensor(
1962 ComplexTensor::new(vec![(-1.0, 0.0), (0.0, 1.0), (0.0, -1.0)], vec![3, 1])
1963 .expect("complex"),
1964 ),
1965 &[],
1966 )
1967 .expect("complex sort")
1968 .into_values();
1969 let Value::ComplexTensor(complex) = complex else {
1970 panic!("expected complex tensor");
1971 };
1972 assert_eq!(
1973 complex.materialize_f64(),
1974 vec![(0.0, -1.0), (0.0, 1.0), (-1.0, 0.0)]
1975 );
1976 let Value::Tensor(complex_indices) = complex_indices else {
1977 panic!("expected complex indices");
1978 };
1979 assert_eq!(complex_indices.materialize_f64(), vec![3.0, 2.0, 1.0]);
1980
1981 let (phase_boundary, phase_boundary_indices) = evaluate(
1982 Value::ComplexTensor(
1983 ComplexTensor::new(vec![(-1.0, 0.0), (-1.0, -0.0)], vec![2, 1])
1984 .expect("phase boundary"),
1985 ),
1986 &[],
1987 )
1988 .expect("phase-boundary sort")
1989 .into_values();
1990 let Value::ComplexTensor(phase_boundary) = phase_boundary else {
1991 panic!("expected complex phase-boundary tensor");
1992 };
1993 assert_eq!(
1994 phase_boundary.materialize_f64(),
1995 vec![(-1.0, 0.0), (-1.0, -0.0)]
1996 );
1997 let Value::Tensor(phase_boundary_indices) = phase_boundary_indices else {
1998 panic!("expected complex phase-boundary indices");
1999 };
2000 assert_eq!(phase_boundary_indices.materialize_f64(), vec![1.0, 2.0]);
2001 }
2002
2003 #[test]
2004 fn sort_preserves_logical_class_and_accepts_every_integer_dimension_class() {
2005 let (logical, logical_indices) = evaluate(
2006 Value::LogicalArray(LogicalArray::new(vec![1, 0, 1], vec![3, 1]).expect("logical")),
2007 &[],
2008 )
2009 .expect("logical sort")
2010 .into_values();
2011 let Value::LogicalArray(logical) = logical else {
2012 panic!("expected logical array");
2013 };
2014 assert_eq!(logical.data, vec![0, 1, 1]);
2015 let Value::Tensor(logical_indices) = logical_indices else {
2016 panic!("expected logical indices");
2017 };
2018 assert_eq!(logical_indices.materialize_f64(), vec![2.0, 1.0, 3.0]);
2019
2020 for dimension in [
2021 IntegerStorage::I8(vec![2]),
2022 IntegerStorage::I16(vec![2]),
2023 IntegerStorage::I32(vec![2]),
2024 IntegerStorage::I64(vec![2]),
2025 IntegerStorage::U8(vec![2]),
2026 IntegerStorage::U16(vec![2]),
2027 IntegerStorage::U32(vec![2]),
2028 IntegerStorage::U64(vec![2]),
2029 ] {
2030 let input = Tensor::new(vec![4.0, 3.0, 2.0, 1.0], vec![2, 2]).expect("input");
2031 let dim = Value::Tensor(Tensor::new_integer(dimension, vec![1, 1]).expect("dimension"));
2032 let sorted = evaluate(Value::Tensor(input), &[dim])
2033 .expect("typed dimension")
2034 .into_sorted_value();
2035 let Value::Tensor(sorted) = sorted else {
2036 panic!("expected tensor");
2037 };
2038 assert_eq!(sorted.materialize_f64(), vec![2.0, 1.0, 4.0, 3.0]);
2039 }
2040 }
2041
2042 #[test]
2043 fn sort_rejects_more_than_two_outputs() {
2044 let _outputs = crate::output_count::push_output_count(Some(3));
2045 let error = sort_builtin(
2046 Value::Tensor(Tensor::new(vec![2.0, 1.0], vec![2, 1]).expect("input")),
2047 Vec::new(),
2048 )
2049 .expect_err("too many outputs");
2050 assert!(error.message().contains("maximum is 2"));
2051 }
2052
2053 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2054 #[test]
2055 fn sort_empty_tensor() {
2056 let tensor = Tensor::new(Vec::new(), vec![0, 3]).unwrap();
2057 let eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("evaluate");
2058 let (sorted, indices) = eval.into_values();
2059 match sorted {
2060 Value::Tensor(t) => {
2061 assert!(t.materialize_f64().is_empty());
2062 assert_eq!(t.shape, tensor.shape);
2063 }
2064 other => panic!("expected tensor, got {other:?}"),
2065 }
2066 match indices {
2067 Value::Tensor(t) => assert!(t.materialize_f64().is_empty()),
2068 other => panic!("expected tensor, got {other:?}"),
2069 }
2070 }
2071
2072 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2073 #[test]
2074 fn sort_dim_greater_than_ndims() {
2075 let tensor = Tensor::new(vec![4.0, 2.0, 3.0, 1.0], vec![2, 2]).unwrap();
2076 let eval = evaluate(
2077 Value::Tensor(tensor.clone()),
2078 &[Value::Int(IntValue::I32(3))],
2079 )
2080 .expect("evaluate");
2081 let (sorted, indices) = eval.into_values();
2082 match sorted {
2083 Value::Tensor(t) => assert_eq!(t.materialize_f64(), tensor.materialize_f64()),
2084 other => panic!("expected tensor, got {other:?}"),
2085 }
2086 match indices {
2087 Value::Tensor(t) => assert!(t
2088 .materialize_f64()
2089 .iter()
2090 .all(|v| (*v - 1.0).abs() < f64::EPSILON)),
2091 other => panic!("expected tensor, got {other:?}"),
2092 }
2093 }
2094
2095 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2096 #[test]
2097 fn sort_supports_missing_placement_and_rejects_unknown_values() {
2098 let input =
2099 Value::Tensor(Tensor::new(vec![2.0, f64::NAN, 1.0], vec![3, 1]).expect("input"));
2100 let first = sort_builtin(
2101 input.clone(),
2102 vec![Value::from("MissingPlacement"), Value::from("first")],
2103 )
2104 .expect("missing first");
2105 let Value::Tensor(first) = first else {
2106 panic!("expected tensor");
2107 };
2108 assert!(first.materialize_f64()[0].is_nan());
2109 assert_eq!(&first.materialize_f64()[1..], &[1.0, 2.0]);
2110
2111 let last_descending = sort_builtin(
2112 input,
2113 vec![
2114 Value::from("descend"),
2115 Value::from("MissingPlacement"),
2116 Value::from("last"),
2117 ],
2118 )
2119 .expect("missing last");
2120 let Value::Tensor(last_descending) = last_descending else {
2121 panic!("expected tensor");
2122 };
2123 assert_eq!(&last_descending.materialize_f64()[..2], &[2.0, 1.0]);
2124 assert!(last_descending.materialize_f64()[2].is_nan());
2125
2126 let err = sort_builtin(
2127 Value::Tensor(Tensor::new(vec![1.0], vec![1, 1]).unwrap()),
2128 vec![Value::from("MissingPlacement"), Value::from("middle")],
2129 )
2130 .unwrap_err();
2131 assert_eq!(
2132 err.identifier(),
2133 SORT_ERROR_MISSINGPLACEMENT_UNKNOWN.identifier
2134 );
2135 }
2136
2137 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2138 #[test]
2139 fn sort_invalid_comparison_method_errors() {
2140 let err = sort_builtin(
2141 Value::Tensor(Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap()),
2142 vec![Value::from("ComparisonMethod"), Value::from("unknown")],
2143 )
2144 .unwrap_err();
2145 assert_eq!(
2146 err.identifier(),
2147 SORT_ERROR_COMPARISON_METHOD_UNKNOWN.identifier
2148 );
2149 }
2150
2151 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2152 #[test]
2153 fn sort_invalid_comparison_method_value_errors() {
2154 let err = sort_builtin(
2155 Value::Tensor(Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap()),
2156 vec![
2157 Value::from("ComparisonMethod"),
2158 Value::Int(IntValue::I32(1)),
2159 ],
2160 )
2161 .unwrap_err();
2162 assert_eq!(
2163 err.identifier(),
2164 SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.identifier
2165 );
2166 }
2167
2168 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2169 #[test]
2170 fn sort_dimension_zero_errors() {
2171 let err = sort_builtin(
2172 Value::Tensor(Tensor::new(vec![1.0], vec![1, 1]).unwrap()),
2173 vec![Value::Num(0.0)],
2174 )
2175 .unwrap_err();
2176 assert_eq!(err.identifier(), SORT_ERROR_INVALID_DIMENSION.identifier);
2177 }
2178
2179 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2180 #[test]
2181 fn sort_gpu_round_trip() {
2182 test_support::with_test_provider(|provider| {
2183 let tensor = Tensor::new(vec![3.0, 1.0, 2.0], vec![3, 1]).unwrap();
2184 let view = runmat_accelerate_api::HostTensorView {
2185 data: &tensor.materialize_f64(),
2186 shape: &tensor.shape,
2187 };
2188 let handle = provider.upload(&view).expect("upload");
2189 let eval = evaluate(Value::GpuTensor(handle), &[]).expect("evaluate");
2190 let (sorted, indices) = eval.into_values();
2191 assert!(matches!(sorted, Value::GpuTensor(_)));
2192 assert!(matches!(indices, Value::GpuTensor(_)));
2193 let sorted = test_support::gather(sorted).expect("gather sorted");
2194 assert_eq!(sorted.materialize_f64(), vec![1.0, 2.0, 3.0]);
2195 let indices = test_support::gather(indices).expect("gather indices");
2196 assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 1.0]);
2197 });
2198 }
2199
2200 #[test]
2201 fn sort_gpu_preserves_exact_wide_integer_and_logical_residency() {
2202 test_support::with_test_provider(|provider| {
2203 let integer = Tensor::new_integer(
2204 IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993]),
2205 vec![3, 1],
2206 )
2207 .expect("integer");
2208 let handle = gpu_helpers::upload_tensor(provider, &integer).expect("upload integer");
2209 let (sorted, indices) = evaluate(Value::GpuTensor(handle), &[])
2210 .expect("integer GPU sort")
2211 .into_values();
2212 assert!(matches!(sorted, Value::GpuTensor(_)));
2213 assert!(matches!(indices, Value::GpuTensor(_)));
2214 let sorted = test_support::gather(sorted).expect("gather integer");
2215 assert_eq!(
2216 sorted.into_numeric_storage().expect("integer storage"),
2217 NumericStorage::U64(vec![0, 9_007_199_254_740_993, u64::MAX])
2218 );
2219
2220 let logical = provider
2221 .upload(&runmat_accelerate_api::HostTensorView {
2222 data: &[1.0, 0.0, 1.0],
2223 shape: &[3, 1],
2224 })
2225 .expect("upload logical");
2226 runmat_accelerate_api::set_handle_logical(&logical, true);
2227 let (sorted, _) = evaluate(Value::GpuTensor(logical), &[])
2228 .expect("logical GPU sort")
2229 .into_values();
2230 let Value::GpuTensor(sorted_handle) = &sorted else {
2231 panic!("expected resident logical output");
2232 };
2233 assert!(runmat_accelerate_api::handle_is_logical(sorted_handle));
2234 let sorted = test_support::gather(sorted).expect("gather logical");
2235 assert_eq!(sorted.materialize_f64(), vec![0.0, 1.0, 1.0]);
2236 });
2237 }
2238
2239 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2240 #[test]
2241 #[cfg(feature = "wgpu")]
2242 fn sort_wgpu_matches_cpu() {
2243 let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2244 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2245 );
2246 let tensor = Tensor::new(vec![4.0, 1.0, 3.0, 2.0], vec![4, 1]).unwrap();
2247 let cpu_eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("cpu sort");
2248 let (cpu_sorted, cpu_indices) = cpu_eval.into_values();
2249
2250 let gpu_view = runmat_accelerate_api::HostTensorView {
2251 data: &tensor.materialize_f64(),
2252 shape: &tensor.shape,
2253 };
2254 let provider = runmat_accelerate_api::provider().expect("wgpu provider");
2255 let handle = provider.upload(&gpu_view).expect("upload");
2256 let gpu_eval = evaluate(Value::GpuTensor(handle), &[]).expect("gpu sort");
2257 let (gpu_sorted, gpu_indices) = gpu_eval.into_values();
2258
2259 let cpu_sorted_tensor = match cpu_sorted {
2260 Value::Tensor(t) => t,
2261 Value::Num(n) => Tensor::new(vec![n], vec![1, 1]).unwrap(),
2262 other => panic!("unexpected CPU sorted value {other:?}"),
2263 };
2264 let cpu_indices_tensor = match cpu_indices {
2265 Value::Tensor(t) => t,
2266 Value::Num(n) => Tensor::new(vec![n], vec![1, 1]).unwrap(),
2267 other => panic!("unexpected CPU indices value {other:?}"),
2268 };
2269 assert!(matches!(gpu_sorted, Value::GpuTensor(_)));
2270 assert!(matches!(gpu_indices, Value::GpuTensor(_)));
2271 let gpu_sorted_tensor = test_support::gather(gpu_sorted).expect("gather GPU sorted");
2272 let gpu_indices_tensor = test_support::gather(gpu_indices).expect("gather GPU indices");
2273
2274 assert_eq!(
2275 gpu_sorted_tensor.materialize_f64(),
2276 cpu_sorted_tensor.materialize_f64()
2277 );
2278 assert_eq!(
2279 gpu_indices_tensor.materialize_f64(),
2280 cpu_indices_tensor.materialize_f64()
2281 );
2282 }
2283
2284 #[test]
2285 #[cfg(feature = "wgpu")]
2286 fn sort_wgpu_abs_ties_follow_phase() {
2287 let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2288 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2289 );
2290 let provider = runmat_accelerate_api::provider().expect("wgpu provider");
2291 let input = provider
2292 .upload(&runmat_accelerate_api::HostTensorView {
2293 data: &[-3.0, 3.0],
2294 shape: &[2, 1],
2295 })
2296 .expect("upload");
2297 let (sorted, indices) = evaluate(
2298 Value::GpuTensor(input),
2299 &[Value::from("ComparisonMethod"), Value::from("abs")],
2300 )
2301 .expect("wgpu sort")
2302 .into_values();
2303 let sorted = test_support::gather(sorted).expect("gather sorted");
2304 assert_eq!(sorted.materialize_f64(), vec![3.0, -3.0]);
2305 let indices = test_support::gather(indices).expect("gather indices");
2306 assert_eq!(indices.materialize_f64(), vec![2.0, 1.0]);
2307 }
2308}