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