1use runmat_accelerate_api::ProviderFindResult;
4use runmat_builtins::{
5 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
6 BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
7 BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
8 BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
9 BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
10 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
11 ResolveContext, Type,
12};
13use runmat_macros::runtime_builtin;
14use runmat_value::{
15 ComplexTensor, IntValue, IntegerComplexStorage, IntegerStorage, LogicalArray, Tensor, Value,
16};
17
18use super::common::fits_positive_platform_index;
19use crate::builtins::array::type_resolvers::column_vector_type;
20use crate::builtins::common::arg_tokens::ArgToken;
21use crate::builtins::common::random_args::complex_tensor_into_value;
22use crate::builtins::common::spec::{
23 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
24 ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
25};
26use crate::builtins::common::{gpu_helpers, tensor};
27use crate::{build_runtime_error, RuntimeError};
28
29#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::array::indexing::find")]
30pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
31 name: "find",
32 op_kind: GpuOpKind::Custom("find"),
33 supported_precisions: &[ScalarType::F64],
34 broadcast: BroadcastSemantics::None,
35 provider_hooks: &[ProviderHook::Custom("find")],
36 constant_strategy: ConstantStrategy::InlineLiteral,
37 residency: ResidencyPolicy::NewHandle,
38 nan_mode: ReductionNaN::Include,
39 two_pass_threshold: None,
40 workgroup_size: None,
41 accepts_nan_mode: false,
42 notes: "Providers execute find directly only when they can return exact f64 indices; f32, logical, and integer cases use a correctness-first host fallback and restore resident outputs.",
43};
44
45#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::array::indexing::find")]
46pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
47 name: "find",
48 shape: ShapeRequirements::Any,
49 constant_strategy: ConstantStrategy::InlineLiteral,
50 elementwise: None,
51 reduction: None,
52 emits_nan: false,
53 notes: "Find drives control flow and currently bypasses fusion; metadata is present for completeness only.",
54};
55
56fn find_type(args: &[Type], _ctx: &ResolveContext) -> Type {
57 if matches!(
58 args.first(),
59 Some(Type::Tensor {
60 shape: Some(shape)
61 }) if shape.len() == 2 && shape.first() == Some(&Some(1))
62 ) {
63 return Type::Tensor {
64 shape: Some(vec![Some(1), None]),
65 };
66 }
67 column_vector_type()
68}
69
70const BUILTIN_NAME: &str = "find";
71
72const FIND_DIRECTION_ONLY_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
73 id: "find-direction-only",
74 mode: BuiltinExtensionMode::RunMatOnly,
75 description: "find(X,direction) is a RunMat convenience extension",
76 error_identifier: Some("RunMat:compatibility:FindDirectionOnlyExtension"),
77};
78const FIND_INTEGER_SPARSE_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
79 id: "find-integer-sparse-input",
80 mode: BuiltinExtensionMode::RunMatOnly,
81 description: "find on typed-integer sparse storage is a RunMat extension",
82 error_identifier: Some("RunMat:compatibility:FindIntegerSparseExtension"),
83};
84pub const FIND_EXTENSIONS: [BuiltinExtensionDescriptor; 2] =
85 [FIND_DIRECTION_ONLY_EXTENSION, FIND_INTEGER_SPARSE_EXTENSION];
86
87const FIND_INTEGER_X_INPUTS: [BuiltinIntegerInputCapability; 1] = [BuiltinIntegerInputCapability {
88 name: "X",
89 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
90 availability: BuiltinIntegerInputAvailability::Documented,
91 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
92 notes: "All eight integer classes use authoritative storage for the exact nonzero predicate.",
93}];
94const FIND_INTEGER_K_INPUTS: [BuiltinIntegerInputCapability; 1] =
95 [BuiltinIntegerInputCapability {
96 name: "K",
97 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
98 availability: BuiltinIntegerInputAvailability::Documented,
99 scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
100 notes: "K is an exact positive scalar count; zero, negative, and out-of-platform-range values reject.",
101 }];
102const FIND_INTEGER_SPARSE_INPUTS: [BuiltinIntegerInputCapability; 1] =
103 [BuiltinIntegerInputCapability {
104 name: "X",
105 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
106 availability: BuiltinIntegerInputAvailability::RunMatOnly,
107 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
108 notes: "MATLAB sparse values are single, double, or logical; typed-integer sparse storage is RunMat-only.",
109 }];
110pub const FIND_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 4] = [
111 BuiltinIntegerCapabilityDescriptor {
112 form: "k = find(integer_X,___)",
113 inputs: &FIND_INTEGER_X_INPUTS,
114 computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
115 output_class: BuiltinIntegerOutputClassRule::Double,
116 overflow: BuiltinIntegerOverflowRule::NotApplicable,
117 backend: BuiltinIntegerBackendRule::GatherFallback,
118 overload: BuiltinIntegerOverloadKind::FunctionSpecific,
119 notes: "Linear indices are exact binary64 indices; resident integer values gather exactly through their owning provider.",
120 },
121 BuiltinIntegerCapabilityDescriptor {
122 form: "[row,col,v] = find(integer_X,___)",
123 inputs: &FIND_INTEGER_X_INPUTS,
124 computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
125 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
126 overflow: BuiltinIntegerOverflowRule::NotApplicable,
127 backend: BuiltinIntegerBackendRule::GatherFallback,
128 overload: BuiltinIntegerOverloadKind::FunctionSpecific,
129 notes: "row and col are exact doubles; v preserves the authoritative integer class and value.",
130 },
131 BuiltinIntegerCapabilityDescriptor {
132 form: "find(X,integer_K[,direction])",
133 inputs: &FIND_INTEGER_K_INPUTS,
134 computation_domain: BuiltinIntegerComputationDomain::Structural,
135 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
136 overflow: BuiltinIntegerOverflowRule::Error,
137 backend: BuiltinIntegerBackendRule::HostOnly,
138 overload: BuiltinIntegerOverloadKind::StructuralParameter,
139 notes: "The positive count is converted exactly to a platform index before any input traversal.",
140 },
141 BuiltinIntegerCapabilityDescriptor {
142 form: "[k|row,col,v] = find(integer_sparse_X,___)",
143 inputs: &FIND_INTEGER_SPARSE_INPUTS,
144 computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
145 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
146 overflow: BuiltinIntegerOverflowRule::NotApplicable,
147 backend: BuiltinIntegerBackendRule::HostOnly,
148 overload: BuiltinIntegerOverloadKind::FunctionSpecific,
149 notes: "Strict compatibility gates this RunMat-only form before CSC traversal; v preserves integer storage.",
150 },
151];
152
153const FIND_OUTPUT_LINEAR: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
154 name: "idx",
155 ty: BuiltinParamType::NumericArray,
156 arity: BuiltinParamArity::Required,
157 default: None,
158 description: "Linear indices of non-zero elements.",
159}];
160
161const FIND_OUTPUT_ROW_COL: [BuiltinParamDescriptor; 2] = [
162 BuiltinParamDescriptor {
163 name: "row",
164 ty: BuiltinParamType::NumericArray,
165 arity: BuiltinParamArity::Required,
166 default: None,
167 description: "Row subscripts of non-zero elements.",
168 },
169 BuiltinParamDescriptor {
170 name: "col",
171 ty: BuiltinParamType::NumericArray,
172 arity: BuiltinParamArity::Required,
173 default: None,
174 description: "Column subscripts of non-zero elements.",
175 },
176];
177
178const FIND_OUTPUT_ROW_COL_VAL: [BuiltinParamDescriptor; 3] = [
179 BuiltinParamDescriptor {
180 name: "row",
181 ty: BuiltinParamType::NumericArray,
182 arity: BuiltinParamArity::Required,
183 default: None,
184 description: "Row subscripts of non-zero elements.",
185 },
186 BuiltinParamDescriptor {
187 name: "col",
188 ty: BuiltinParamType::NumericArray,
189 arity: BuiltinParamArity::Required,
190 default: None,
191 description: "Column subscripts of non-zero elements.",
192 },
193 BuiltinParamDescriptor {
194 name: "v",
195 ty: BuiltinParamType::Any,
196 arity: BuiltinParamArity::Required,
197 default: None,
198 description: "Values at the reported row/column locations.",
199 },
200];
201
202const FIND_INPUTS_BASE: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
203 name: "X",
204 ty: BuiltinParamType::Any,
205 arity: BuiltinParamArity::Required,
206 default: None,
207 description: "Input array to search.",
208}];
209
210const FIND_INPUTS_LIMIT: [BuiltinParamDescriptor; 2] = [
211 BuiltinParamDescriptor {
212 name: "X",
213 ty: BuiltinParamType::Any,
214 arity: BuiltinParamArity::Required,
215 default: None,
216 description: "Input array to search.",
217 },
218 BuiltinParamDescriptor {
219 name: "K",
220 ty: BuiltinParamType::NumericScalar,
221 arity: BuiltinParamArity::Required,
222 default: None,
223 description: "Maximum number of indices to return.",
224 },
225];
226
227const FIND_INPUTS_LIMIT_DIR: [BuiltinParamDescriptor; 3] = [
228 BuiltinParamDescriptor {
229 name: "X",
230 ty: BuiltinParamType::Any,
231 arity: BuiltinParamArity::Required,
232 default: None,
233 description: "Input array to search.",
234 },
235 BuiltinParamDescriptor {
236 name: "K",
237 ty: BuiltinParamType::NumericScalar,
238 arity: BuiltinParamArity::Required,
239 default: None,
240 description: "Maximum number of indices to return.",
241 },
242 BuiltinParamDescriptor {
243 name: "direction",
244 ty: BuiltinParamType::StringScalar,
245 arity: BuiltinParamArity::Required,
246 default: Some("\"first\""),
247 description: "Direction selector: `\"first\"` or `\"last\"`.",
248 },
249];
250
251const FIND_SIGNATURES: [BuiltinSignatureDescriptor; 7] = [
252 BuiltinSignatureDescriptor {
253 label: "idx = find(X)",
254 inputs: &FIND_INPUTS_BASE,
255 outputs: &FIND_OUTPUT_LINEAR,
256 },
257 BuiltinSignatureDescriptor {
258 label: "idx = find(X, K)",
259 inputs: &FIND_INPUTS_LIMIT,
260 outputs: &FIND_OUTPUT_LINEAR,
261 },
262 BuiltinSignatureDescriptor {
263 label: "idx = find(X, K, direction)",
264 inputs: &FIND_INPUTS_LIMIT_DIR,
265 outputs: &FIND_OUTPUT_LINEAR,
266 },
267 BuiltinSignatureDescriptor {
268 label: "[row, col] = find(X)",
269 inputs: &FIND_INPUTS_BASE,
270 outputs: &FIND_OUTPUT_ROW_COL,
271 },
272 BuiltinSignatureDescriptor {
273 label: "[row, col] = find(X, K, direction)",
274 inputs: &FIND_INPUTS_LIMIT_DIR,
275 outputs: &FIND_OUTPUT_ROW_COL,
276 },
277 BuiltinSignatureDescriptor {
278 label: "[row, col, v] = find(X)",
279 inputs: &FIND_INPUTS_BASE,
280 outputs: &FIND_OUTPUT_ROW_COL_VAL,
281 },
282 BuiltinSignatureDescriptor {
283 label: "[row, col, v] = find(X, K, direction)",
284 inputs: &FIND_INPUTS_LIMIT_DIR,
285 outputs: &FIND_OUTPUT_ROW_COL_VAL,
286 },
287];
288
289const FIND_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
290 code: "RM.FIND.INVALID_INPUT",
291 identifier: Some("RunMat:find:InvalidInput"),
292 when: "Input type or option arguments are not valid for find.",
293 message: "find: invalid input arguments",
294};
295
296const FIND_ERROR_PROVIDER_OUTPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
297 code: "RM.FIND.PROVIDER_OUTPUT",
298 identifier: Some("RunMat:find:ProviderOutput"),
299 when: "GPU provider does not return expected output buffers for requested nargout.",
300 message: "find: provider output buffer mismatch",
301};
302
303const FIND_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
304 code: "RM.FIND.INTERNAL",
305 identifier: Some("RunMat:find:InternalError"),
306 when: "Internal tensor conversion/materialization fails while building outputs.",
307 message: "find: internal error",
308};
309
310const FIND_ERRORS: [BuiltinErrorDescriptor; 3] = [
311 FIND_ERROR_INVALID_INPUT,
312 FIND_ERROR_PROVIDER_OUTPUT,
313 FIND_ERROR_INTERNAL,
314];
315
316pub const FIND_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
317 signatures: &FIND_SIGNATURES,
318 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
319 completion_policy: BuiltinCompletionPolicy::Public,
320 errors: &FIND_ERRORS,
321};
322
323fn find_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
324 find_error_with_message(error.message, error)
325}
326
327fn find_error_with_message(
328 message: impl Into<String>,
329 error: &'static BuiltinErrorDescriptor,
330) -> RuntimeError {
331 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
332 if let Some(identifier) = error.identifier {
333 builder = builder.with_identifier(identifier);
334 }
335 builder.build()
336}
337
338fn parse_find_tokens(tokens: &[ArgToken]) -> crate::BuiltinResult<FindOptions> {
339 match tokens.len() {
340 0 => Ok(FindOptions::default()),
341 1 => {
342 if let Some(direction) = token_to_direction(&tokens[0])? {
343 let limit = if matches!(direction, FindDirection::Last) {
344 Some(1)
345 } else {
346 None
347 };
348 Ok(FindOptions { limit, direction })
349 } else {
350 let limit = token_to_limit(&tokens[0])?;
351 Ok(FindOptions {
352 limit: Some(limit),
353 direction: FindDirection::First,
354 })
355 }
356 }
357 2 => {
358 let limit = token_to_limit(&tokens[0])?;
359 let direction = token_to_direction(&tokens[1])?.ok_or_else(|| {
360 find_error_with_message(
361 "find: third argument must be 'first' or 'last'",
362 &FIND_ERROR_INVALID_INPUT,
363 )
364 })?;
365 Ok(FindOptions {
366 limit: Some(limit),
367 direction,
368 })
369 }
370 _ => Err(find_error_with_message(
371 "find: too many input arguments",
372 &FIND_ERROR_INVALID_INPUT,
373 )),
374 }
375}
376
377fn token_to_direction(token: &ArgToken) -> crate::BuiltinResult<Option<FindDirection>> {
378 match token {
379 ArgToken::String(text) => match text.as_str() {
380 "first" => Ok(Some(FindDirection::First)),
381 "last" => Ok(Some(FindDirection::Last)),
382 _ => Err(find_error_with_message(
383 "find: direction must be 'first' or 'last'",
384 &FIND_ERROR_INVALID_INPUT,
385 )),
386 },
387 _ => Ok(None),
388 }
389}
390
391fn token_to_limit(token: &ArgToken) -> crate::BuiltinResult<usize> {
392 match token {
393 ArgToken::Number(value) => parse_limit_scalar(*value),
394 ArgToken::Integer(value) => parse_limit_integer(value),
395 _ => Err(find_error_with_message(
396 "find: second argument must be a scalar",
397 &FIND_ERROR_INVALID_INPUT,
398 )),
399 }
400}
401
402#[runtime_builtin(
403 name = "find",
404 category = "array/indexing",
405 summary = "Locate nonzero indices and values.",
406 keywords = "find,nonzero,indices,row,column,gpu",
407 accel = "custom",
408 type_resolver(find_type),
409 descriptor(crate::builtins::array::indexing::find::FIND_DESCRIPTOR),
410 extensions(crate::builtins::array::indexing::find::FIND_EXTENSIONS),
411 integer_capabilities(crate::builtins::array::indexing::find::FIND_INTEGER_CAPABILITIES),
412 builtin_path = "crate::builtins::array::indexing::find"
413)]
414async fn find_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
415 let eval = evaluate(value, &rest).await?;
416 if let Some(out_count) = crate::output_count::current_output_count() {
417 if out_count == 0 {
418 return Ok(Value::OutputList(Vec::new()));
419 }
420 if out_count <= 1 {
421 let linear = eval.linear_value()?;
422 return Ok(crate::output_count::output_list_with_padding(
423 out_count,
424 vec![linear],
425 ));
426 }
427 let rows = eval.row_value()?;
428 let cols = eval.column_value()?;
429 let mut outputs = vec![rows, cols];
430 if out_count >= 3 {
431 outputs.push(eval.values_value()?);
432 }
433 return Ok(crate::output_count::output_list_with_padding(
434 out_count, outputs,
435 ));
436 }
437 eval.linear_value()
438}
439
440pub async fn evaluate(value: Value, args: &[Value]) -> crate::BuiltinResult<FindEval> {
442 if args.len() == 1
443 && matches!(
444 crate::builtins::common::arg_tokens::tokens_from_values(args).first(),
445 Some(ArgToken::String(_))
446 )
447 {
448 crate::compatibility::ensure_builtin_extension_enabled(
449 &FIND_DIRECTION_ONLY_EXTENSION,
450 BUILTIN_NAME,
451 )?;
452 }
453 if matches!(&value, Value::SparseTensor(sparse) if sparse.integer_storage().is_some()) {
454 crate::compatibility::ensure_builtin_extension_enabled(
455 &FIND_INTEGER_SPARSE_EXTENSION,
456 BUILTIN_NAME,
457 )?;
458 }
459 let options = parse_options(args).await?;
460 match value {
461 Value::GpuTensor(handle) => {
462 let owner = runmat_accelerate_api::provider_for_handle(&handle).ok_or_else(|| {
463 find_error_with_message(
464 "find: no acceleration provider owns the input handle",
465 &FIND_ERROR_INTERNAL,
466 )
467 })?;
468 let provider_has_exact_double_indices = matches!(
469 runmat_accelerate_api::handle_precision(&handle),
470 Some(runmat_accelerate_api::ProviderPrecision::F64)
471 );
472 let provider_indices_are_exact_double = provider_has_exact_double_indices
473 && !runmat_accelerate_api::handle_is_logical(&handle)
474 && runmat_accelerate_api::handle_integer_type(&handle).is_none();
475 if provider_indices_are_exact_double {
476 if let Some(result) = try_provider_find(owner, &handle, &options) {
477 return Ok(FindEval::from_gpu(result));
478 }
479 }
480 let (storage, _) = materialize_input(Value::GpuTensor(handle)).await?;
481 let result = compute_find(&storage, &options);
482 Ok(FindEval::from_host(result, Some(owner)))
483 }
484 Value::SparseTensor(sparse) => {
485 let result = compute_find_sparse(&sparse, &options);
486 Ok(FindEval::from_host(result, None))
487 }
488 other => {
489 let (storage, _) = materialize_input(other).await?;
490 let result = compute_find(&storage, &options);
491 Ok(FindEval::from_host(result, None))
492 }
493 }
494}
495
496fn try_provider_find(
497 provider: &'static dyn runmat_accelerate_api::AccelProvider,
498 handle: &runmat_accelerate_api::GpuTensorHandle,
499 options: &FindOptions,
500) -> Option<ProviderFindResult> {
501 if matches!(options.direction, FindDirection::Last) {
502 return None;
503 }
504 let direction = match options.direction {
505 FindDirection::First => runmat_accelerate_api::FindDirection::First,
506 FindDirection::Last => runmat_accelerate_api::FindDirection::Last,
507 };
508 let limit = options.effective_limit();
509 let mut result = provider.find(handle, limit, direction).ok()?;
510 if is_row_vector_shape(&handle.shape) {
511 result.linear.shape = vec![1, result.linear.shape.first().copied().unwrap_or(0)];
512 }
513 Some(result)
514}
515
516#[derive(Debug, Clone, Copy, PartialEq, Eq)]
517enum FindDirection {
518 First,
519 Last,
520}
521
522#[derive(Debug, Clone)]
523struct FindOptions {
524 limit: Option<usize>,
525 direction: FindDirection,
526}
527
528impl Default for FindOptions {
529 fn default() -> Self {
530 Self {
531 limit: None,
532 direction: FindDirection::First,
533 }
534 }
535}
536
537impl FindOptions {
538 fn effective_limit(&self) -> Option<usize> {
539 match self.direction {
540 FindDirection::Last => self.limit.or(Some(1)),
541 FindDirection::First => self.limit,
542 }
543 }
544}
545
546#[derive(Clone)]
547enum DataStorage {
548 Real(Tensor),
549 Logical(LogicalArray),
550 Complex(ComplexTensor),
551}
552
553impl DataStorage {
554 fn shape(&self) -> &[usize] {
555 match self {
556 DataStorage::Real(t) => &t.shape,
557 DataStorage::Logical(t) => &t.shape,
558 DataStorage::Complex(t) => &t.shape,
559 }
560 }
561}
562
563#[derive(Clone)]
564struct FindResult {
565 shape: Vec<usize>,
566 indices: Vec<usize>,
567 values: FindValues,
568}
569
570#[derive(Clone)]
571enum FindValues {
572 Real(Vec<f64>),
573 F32(Vec<f32>),
574 Logical(Vec<u8>),
575 Integer(IntegerStorage),
576 Complex(Vec<(f64, f64)>),
577 IntegerComplex(IntegerComplexStorage),
578}
579
580pub struct FindEval {
581 inner: FindEvalInner,
582}
583
584enum FindEvalInner {
585 Host {
586 result: FindResult,
587 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
588 },
589 Gpu {
590 result: ProviderFindResult,
591 },
592}
593
594impl FindEval {
595 fn from_host(
596 result: FindResult,
597 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
598 ) -> Self {
599 Self {
600 inner: FindEvalInner::Host {
601 result,
602 output_provider,
603 },
604 }
605 }
606
607 fn from_gpu(result: ProviderFindResult) -> Self {
608 Self {
609 inner: FindEvalInner::Gpu { result },
610 }
611 }
612
613 pub fn linear_value(&self) -> crate::BuiltinResult<Value> {
614 match &self.inner {
615 FindEvalInner::Host {
616 result,
617 output_provider,
618 } => {
619 let tensor = result.linear_tensor()?;
620 Ok(tensor_to_value(tensor, *output_provider))
621 }
622 FindEvalInner::Gpu { result } => Ok(Value::GpuTensor(result.linear.clone())),
623 }
624 }
625
626 pub fn row_value(&self) -> crate::BuiltinResult<Value> {
627 match &self.inner {
628 FindEvalInner::Host {
629 result,
630 output_provider,
631 } => {
632 let tensor = result.row_tensor()?;
633 Ok(tensor_to_value(tensor, *output_provider))
634 }
635 FindEvalInner::Gpu { result } => Ok(Value::GpuTensor(result.rows.clone())),
636 }
637 }
638
639 pub fn column_value(&self) -> crate::BuiltinResult<Value> {
640 match &self.inner {
641 FindEvalInner::Host {
642 result,
643 output_provider,
644 } => {
645 let tensor = result.column_tensor()?;
646 Ok(tensor_to_value(tensor, *output_provider))
647 }
648 FindEvalInner::Gpu { result } => Ok(Value::GpuTensor(result.cols.clone())),
649 }
650 }
651
652 pub fn values_value(&self) -> crate::BuiltinResult<Value> {
653 match &self.inner {
654 FindEvalInner::Host {
655 result,
656 output_provider,
657 } => result.values_value(*output_provider),
658 FindEvalInner::Gpu { result } => result
659 .values
660 .as_ref()
661 .map(|handle| Value::GpuTensor(handle.clone()))
662 .ok_or_else(|| find_error(&FIND_ERROR_PROVIDER_OUTPUT)),
663 }
664 }
665}
666
667async fn parse_options(args: &[Value]) -> crate::BuiltinResult<FindOptions> {
668 parse_find_tokens(&crate::builtins::common::arg_tokens::tokens_from_values(
669 args,
670 ))
671}
672
673fn parse_limit_integer(value: &IntValue) -> crate::BuiltinResult<usize> {
674 let value = value.try_to_usize().ok_or_else(|| {
675 find_error_with_message(
676 "find: K must be a positive integer within the supported range",
677 &FIND_ERROR_INVALID_INPUT,
678 )
679 })?;
680 if value == 0 {
681 return Err(find_error_with_message(
682 "find: K must be a positive integer",
683 &FIND_ERROR_INVALID_INPUT,
684 ));
685 }
686 Ok(value)
687}
688
689fn parse_limit_scalar(value: f64) -> crate::BuiltinResult<usize> {
690 if !value.is_finite() {
691 return Err(find_error_with_message(
692 "find: K must be a finite, non-negative integer",
693 &FIND_ERROR_INVALID_INPUT,
694 ));
695 }
696 let rounded = value.round();
697 if (rounded - value).abs() > f64::EPSILON {
698 return Err(find_error_with_message(
699 "find: K must be a finite, non-negative integer",
700 &FIND_ERROR_INVALID_INPUT,
701 ));
702 }
703 if rounded <= 0.0 {
704 return Err(find_error_with_message(
705 "find: K must be a positive integer",
706 &FIND_ERROR_INVALID_INPUT,
707 ));
708 }
709 if !fits_positive_platform_index(rounded) {
710 return Err(find_error_with_message(
711 "find: K exceeds the maximum supported index range",
712 &FIND_ERROR_INVALID_INPUT,
713 ));
714 }
715 Ok(rounded as usize)
716}
717
718async fn materialize_input(value: Value) -> crate::BuiltinResult<(DataStorage, bool)> {
719 match value {
720 Value::GpuTensor(handle) => {
721 let is_logical = runmat_accelerate_api::handle_is_logical(&handle);
722 let tensor = gpu_helpers::gather_tensor_async(&handle).await?;
723 if is_logical {
724 let data = (0..tensor::tensor_element_len(&tensor))
725 .map(|index| u8::from(tensor::tensor_value_f64(&tensor, index) != 0.0))
726 .collect();
727 let shape = tensor.shape.clone();
728 return LogicalArray::new(data, shape)
729 .map(|logical| (DataStorage::Logical(logical), true))
730 .map_err(|e| {
731 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
732 });
733 }
734 Ok((DataStorage::Real(tensor), true))
735 }
736 Value::Tensor(tensor) => Ok((DataStorage::Real(tensor), false)),
737 Value::SparseTensor(sparse) => {
738 let dense = if sparse.is_logical() {
739 tensor::logical_to_tensor(&sparse.to_dense_logical().map_err(|e| {
740 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
741 })?)
742 .map_err(|message| find_error_with_message(message, &FIND_ERROR_INTERNAL))?
743 } else {
744 sparse.to_dense().map_err(|e| {
745 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
746 })?
747 };
748 Ok((DataStorage::Real(dense), false))
749 }
750 Value::LogicalArray(logical) => Ok((DataStorage::Logical(logical), false)),
751 Value::Num(n) => {
752 let tensor = Tensor::new(vec![n], vec![1, 1])
753 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))?;
754 Ok((DataStorage::Real(tensor), false))
755 }
756 Value::Int(i) => {
757 let tensor = Tensor::new_integer(integer_storage_from_scalar(&i), vec![1, 1])
758 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))?;
759 Ok((DataStorage::Real(tensor), false))
760 }
761 Value::Bool(b) => LogicalArray::new(vec![u8::from(b)], vec![1, 1])
762 .map(|logical| (DataStorage::Logical(logical), false))
763 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)),
764 Value::Complex(re, im) => {
765 let tensor = ComplexTensor::new(vec![(re, im)], vec![1, 1])
766 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))?;
767 Ok((DataStorage::Complex(tensor), false))
768 }
769 Value::ComplexTensor(tensor) => Ok((DataStorage::Complex(tensor), false)),
770 Value::CharArray(chars) => {
771 let mut data = Vec::with_capacity(chars.data.len());
772 for c in 0..chars.cols {
773 for r in 0..chars.rows {
774 let ch = chars.data[r * chars.cols + c] as u32;
775 data.push(ch as f64);
776 }
777 }
778 let tensor = Tensor::new(data, vec![chars.rows, chars.cols])
779 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))?;
780 Ok((DataStorage::Real(tensor), false))
781 }
782 other => Err(find_error_with_message(
783 format!(
784 "find: unsupported input type {:?}; expected numeric, logical, or char data",
785 other
786 ),
787 &FIND_ERROR_INVALID_INPUT,
788 )),
789 }
790}
791
792fn compute_find(storage: &DataStorage, options: &FindOptions) -> FindResult {
793 let shape = storage.shape().to_vec();
794 let limit = options.effective_limit();
795
796 match storage {
797 DataStorage::Real(tensor) => {
798 let mut indices = Vec::new();
799 let typed_storage = tensor.integer_storage();
800
801 if matches!(limit, Some(0)) {
802 return FindResult::new(shape, indices, find_values_for_tensor(tensor, &[]));
803 }
804
805 let len = typed_storage
806 .map(|storage| storage.len())
807 .unwrap_or_else(|| tensor::tensor_element_len(tensor));
808 match options.direction {
809 FindDirection::First => {
810 for idx in 0..len {
811 let nonzero = typed_storage.map_or_else(
812 || tensor::tensor_value_f64(tensor, idx) != 0.0,
813 |storage| {
814 storage
815 .value_at(idx)
816 .map(|value| !value.is_zero())
817 .expect("typed integer storage is structurally valid")
818 },
819 );
820 if nonzero {
821 indices.push(idx + 1);
822 if limit.is_some_and(|k| indices.len() >= k) {
823 break;
824 }
825 }
826 }
827 }
828 FindDirection::Last => {
829 for idx in (0..len).rev() {
830 let nonzero = typed_storage.map_or_else(
831 || tensor::tensor_value_f64(tensor, idx) != 0.0,
832 |storage| {
833 storage
834 .value_at(idx)
835 .map(|value| !value.is_zero())
836 .expect("typed integer storage is structurally valid")
837 },
838 );
839 if nonzero {
840 indices.push(idx + 1);
841 if limit.is_some_and(|k| indices.len() >= k) {
842 break;
843 }
844 }
845 }
846 }
847 }
848
849 if matches!(options.direction, FindDirection::Last) {
850 indices.reverse();
851 }
852 let values = find_values_for_tensor(tensor, &indices);
853 FindResult::new(shape, indices, values)
854 }
855 DataStorage::Logical(logical) => {
856 let mut indices = Vec::new();
857 if !matches!(options.effective_limit(), Some(0)) {
858 let iter: Box<dyn Iterator<Item = usize>> = match options.direction {
859 FindDirection::First => Box::new(0..logical.data.len()),
860 FindDirection::Last => Box::new((0..logical.data.len()).rev()),
861 };
862 for idx in iter {
863 if logical.data[idx] != 0 {
864 indices.push(idx + 1);
865 if options
866 .effective_limit()
867 .is_some_and(|limit| indices.len() >= limit)
868 {
869 break;
870 }
871 }
872 }
873 }
874 if matches!(options.direction, FindDirection::Last) {
875 indices.reverse();
876 }
877 let values = FindValues::Logical(vec![1; indices.len()]);
878 FindResult::new(shape, indices, values)
879 }
880 DataStorage::Complex(tensor) => {
881 let mut indices = Vec::new();
882 let mut values = Vec::new();
883 let typed_storage = tensor.integer_storage();
884
885 if matches!(limit, Some(0)) {
886 let values = find_values_for_complex_tensor(tensor, &indices, values);
887 return FindResult::new(shape, indices, values);
888 }
889
890 let len = typed_storage
891 .map(|storage| storage.len())
892 .unwrap_or(tensor.materialize_f64().len());
893 match options.direction {
894 FindDirection::First => {
895 for idx in 0..len {
896 let nonzero = typed_storage.map_or_else(
897 || {
898 let (re, im) = tensor.materialize_f64()[idx];
899 re != 0.0 || im != 0.0
900 },
901 |storage| {
902 storage
903 .is_nonzero_at(idx)
904 .expect("typed complex integer storage is structurally valid")
905 },
906 );
907 if nonzero {
908 indices.push(idx + 1);
909 if typed_storage.is_none() {
910 values.push(tensor.materialize_f64()[idx]);
911 }
912 if limit.is_some_and(|k| indices.len() >= k) {
913 break;
914 }
915 }
916 }
917 }
918 FindDirection::Last => {
919 for idx in (0..len).rev() {
920 let nonzero = typed_storage.map_or_else(
921 || {
922 let (re, im) = tensor.materialize_f64()[idx];
923 re != 0.0 || im != 0.0
924 },
925 |storage| {
926 storage
927 .is_nonzero_at(idx)
928 .expect("typed complex integer storage is structurally valid")
929 },
930 );
931 if nonzero {
932 indices.push(idx + 1);
933 if typed_storage.is_none() {
934 values.push(tensor.materialize_f64()[idx]);
935 }
936 if limit.is_some_and(|k| indices.len() >= k) {
937 break;
938 }
939 }
940 }
941 }
942 }
943
944 if matches!(options.direction, FindDirection::Last) {
945 indices.reverse();
946 values.reverse();
947 }
948 let values = find_values_for_complex_tensor(tensor, &indices, values);
949 FindResult::new(shape, indices, values)
950 }
951 }
952}
953
954fn sparse_find_values(
955 sparse: &runmat_value::SparseTensor,
956 real_values: Vec<f64>,
957 single_values: Vec<f32>,
958 logical_values: Vec<u8>,
959 integer_value_indices: &[usize],
960) -> FindValues {
961 if sparse.is_logical() {
962 FindValues::Logical(logical_values)
963 } else if let Some(storage) = sparse.integer_storage() {
964 FindValues::Integer(select_integer_values(storage, integer_value_indices))
965 } else if sparse.as_f32_slice().is_some() {
966 FindValues::F32(single_values)
967 } else {
968 FindValues::Real(real_values)
969 }
970}
971
972fn sparse_stored_value_is_nonzero(sparse: &runmat_value::SparseTensor, index: usize) -> bool {
973 !sparse
974 .numeric_value_at(index)
975 .expect("SparseTensor value storage is consistent")
976 .is_zero()
977}
978
979fn compute_find_sparse(sparse: &runmat_value::SparseTensor, options: &FindOptions) -> FindResult {
980 let shape = vec![sparse.rows, sparse.cols];
981 let limit = options.effective_limit();
982
983 let mut indices = Vec::new();
984 let mut values = Vec::new();
985 let mut single_values = Vec::new();
986 let mut logical_values = Vec::new();
987 let integer_storage = sparse.integer_storage();
988 let floating_values = sparse.as_f64_slice();
989 let native_single_values = sparse.as_f32_slice();
990 let mut integer_value_indices = Vec::new();
991
992 if matches!(limit, Some(0)) {
993 let values = sparse_find_values(
994 sparse,
995 values,
996 single_values,
997 logical_values,
998 &integer_value_indices,
999 );
1000 return FindResult::new(shape, indices, values);
1001 }
1002
1003 match options.direction {
1004 FindDirection::First => {
1005 for col in 0..sparse.cols {
1006 let col_start = sparse.col_ptrs[col];
1007 let col_end = sparse.col_ptrs[col + 1];
1008 for idx in col_start..col_end {
1009 let row = sparse.row_indices[idx];
1010 if sparse_stored_value_is_nonzero(sparse, idx) {
1011 let linear_idx = row + col * sparse.rows;
1012 indices.push(linear_idx + 1);
1013 if sparse.is_logical() {
1014 logical_values.push(1);
1015 } else if integer_storage.is_some() {
1016 integer_value_indices.push(idx);
1017 } else if let Some(native_single_values) = native_single_values {
1018 single_values.push(native_single_values[idx]);
1019 } else {
1020 values.push(floating_values.expect("double sparse storage")[idx]);
1021 }
1022 if limit.is_some_and(|k| indices.len() >= k) {
1023 let values = sparse_find_values(
1024 sparse,
1025 values,
1026 single_values,
1027 logical_values,
1028 &integer_value_indices,
1029 );
1030 return FindResult::new(shape, indices, values);
1031 }
1032 }
1033 }
1034 }
1035 }
1036 FindDirection::Last => {
1037 for col in (0..sparse.cols).rev() {
1038 let col_start = sparse.col_ptrs[col];
1039 let col_end = sparse.col_ptrs[col + 1];
1040 for idx in (col_start..col_end).rev() {
1041 let row = sparse.row_indices[idx];
1042 if sparse_stored_value_is_nonzero(sparse, idx) {
1043 let linear_idx = row + col * sparse.rows;
1044 indices.push(linear_idx + 1);
1045 if sparse.is_logical() {
1046 logical_values.push(1);
1047 } else if integer_storage.is_some() {
1048 integer_value_indices.push(idx);
1049 } else if let Some(native_single_values) = native_single_values {
1050 single_values.push(native_single_values[idx]);
1051 } else {
1052 values.push(floating_values.expect("double sparse storage")[idx]);
1053 }
1054 if limit.is_some_and(|k| indices.len() >= k) {
1055 indices.reverse();
1056 values.reverse();
1057 single_values.reverse();
1058 logical_values.reverse();
1059 integer_value_indices.reverse();
1060 let values = sparse_find_values(
1061 sparse,
1062 values,
1063 single_values,
1064 logical_values,
1065 &integer_value_indices,
1066 );
1067 return FindResult::new(shape, indices, values);
1068 }
1069 }
1070 }
1071 }
1072 }
1073 }
1074
1075 if matches!(options.direction, FindDirection::Last) {
1076 indices.reverse();
1077 values.reverse();
1078 single_values.reverse();
1079 logical_values.reverse();
1080 integer_value_indices.reverse();
1081 }
1082 let values = sparse_find_values(
1083 sparse,
1084 values,
1085 single_values,
1086 logical_values,
1087 &integer_value_indices,
1088 );
1089 FindResult::new(shape, indices, values)
1090}
1091
1092fn is_row_vector_shape(shape: &[usize]) -> bool {
1093 shape.len() == 2 && shape.first() == Some(&1)
1094}
1095
1096impl FindResult {
1097 fn new(shape: Vec<usize>, indices: Vec<usize>, values: FindValues) -> Self {
1098 Self {
1099 shape,
1100 indices,
1101 values,
1102 }
1103 }
1104
1105 fn linear_tensor(&self) -> crate::BuiltinResult<Tensor> {
1106 let data = self
1107 .indices
1108 .iter()
1109 .map(|&idx| exact_index_as_f64(idx))
1110 .collect::<crate::BuiltinResult<Vec<_>>>()?;
1111 let shape = if data.is_empty() && matches!(self.shape.as_slice(), [0, 0] | [1, 1]) {
1112 vec![0, 0]
1113 } else if is_row_vector_shape(&self.shape) {
1114 vec![1, data.len()]
1115 } else {
1116 vec![data.len(), 1]
1117 };
1118 Tensor::new(data, shape)
1119 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))
1120 }
1121
1122 fn row_tensor(&self) -> crate::BuiltinResult<Tensor> {
1123 let mut data = Vec::with_capacity(self.indices.len());
1124 let rows = self.shape.first().copied().unwrap_or(1).max(1);
1125 for &idx in &self.indices {
1126 let zero_based = idx - 1;
1127 let row = (zero_based % rows) + 1;
1128 data.push(exact_index_as_f64(row)?);
1129 }
1130 Tensor::new(data, vec![self.indices.len(), 1])
1131 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))
1132 }
1133
1134 fn column_tensor(&self) -> crate::BuiltinResult<Tensor> {
1135 let mut data = Vec::with_capacity(self.indices.len());
1136 let rows = self.shape.first().copied().unwrap_or(1).max(1);
1137 for &idx in &self.indices {
1138 let zero_based = idx - 1;
1139 let col = (zero_based / rows) + 1;
1140 data.push(exact_index_as_f64(col)?);
1141 }
1142 Tensor::new(data, vec![self.indices.len(), 1])
1143 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))
1144 }
1145
1146 fn values_value(
1147 &self,
1148 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
1149 ) -> crate::BuiltinResult<Value> {
1150 match &self.values {
1151 FindValues::Real(values) => {
1152 let tensor = Tensor::new(values.clone(), vec![values.len(), 1]).map_err(|e| {
1153 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1154 })?;
1155 Ok(tensor_to_value(tensor, output_provider))
1156 }
1157 FindValues::F32(values) => {
1158 let tensor =
1159 Tensor::from_f32(values.clone(), vec![values.len(), 1]).map_err(|e| {
1160 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1161 })?;
1162 Ok(tensor_to_value(tensor, output_provider))
1163 }
1164 FindValues::Logical(values) => {
1165 let logical =
1166 LogicalArray::new(values.clone(), vec![values.len(), 1]).map_err(|e| {
1167 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1168 })?;
1169 if let Some(provider) = output_provider {
1170 let tensor = Tensor::new(
1171 values.iter().map(|&value| f64::from(value)).collect(),
1172 logical.shape.clone(),
1173 )
1174 .map_err(|e| {
1175 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1176 })?;
1177 if let Ok(handle) = gpu_helpers::upload_tensor(provider, &tensor) {
1178 return Ok(gpu_helpers::logical_gpu_value(handle));
1179 }
1180 }
1181 Ok(Value::LogicalArray(logical))
1182 }
1183 FindValues::Integer(values) => integer_values_to_value(values.clone(), output_provider),
1184 FindValues::Complex(values) => {
1185 let tensor =
1186 ComplexTensor::new(values.clone(), vec![values.len(), 1]).map_err(|e| {
1187 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1188 })?;
1189 Ok(complex_tensor_to_value(tensor, output_provider))
1190 }
1191 FindValues::IntegerComplex(storage) => {
1192 let tensor = ComplexTensor::new_integer(storage.clone(), vec![storage.len(), 1])
1193 .map_err(|e| {
1194 find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL)
1195 })?;
1196 Ok(complex_tensor_to_value(tensor, output_provider))
1197 }
1198 }
1199 }
1200}
1201
1202fn exact_index_as_f64(index: usize) -> crate::BuiltinResult<f64> {
1203 const MAX_EXACT_BINARY64_INTEGER: u128 = 1_u128 << 53;
1204 if (index as u128) > MAX_EXACT_BINARY64_INTEGER {
1205 return Err(find_error_with_message(
1206 "find: index exceeds the exact binary64 index range",
1207 &FIND_ERROR_INVALID_INPUT,
1208 ));
1209 }
1210 Ok(index as f64)
1211}
1212
1213fn find_values_for_tensor(tensor: &Tensor, indices: &[usize]) -> FindValues {
1214 let Some(storage) = tensor.integer_storage() else {
1215 if let Some(values) = tensor.as_f32_slice() {
1216 return FindValues::F32(indices.iter().map(|index| values[index - 1]).collect());
1217 }
1218 return FindValues::Real(
1219 indices
1220 .iter()
1221 .map(|index| tensor::tensor_value_f64(tensor, index - 1))
1222 .collect(),
1223 );
1224 };
1225 let selected: Vec<usize> = indices.iter().map(|index| index - 1).collect();
1226 FindValues::Integer(select_integer_values(storage, &selected))
1227}
1228
1229fn find_values_for_complex_tensor(
1230 tensor: &ComplexTensor,
1231 indices: &[usize],
1232 values: Vec<(f64, f64)>,
1233) -> FindValues {
1234 let Some(storage) = tensor.integer_storage() else {
1235 return FindValues::Complex(values);
1236 };
1237 let selected: Vec<usize> = indices.iter().map(|index| index - 1).collect();
1238 let real = select_integer_values(&storage.real, &selected);
1239 let imag = select_integer_values(&storage.imag, &selected);
1240 let storage = IntegerComplexStorage::new(real, imag)
1241 .expect("paired typed complex storage preserves class and length through find");
1242 FindValues::IntegerComplex(storage)
1243}
1244
1245fn select_integer_values(storage: &IntegerStorage, indices: &[usize]) -> IntegerStorage {
1246 macro_rules! select {
1247 ($values:expr, $variant:ident) => {
1248 IntegerStorage::$variant(indices.iter().map(|&index| $values[index]).collect())
1249 };
1250 }
1251 match storage {
1252 IntegerStorage::I8(values) => select!(values, I8),
1253 IntegerStorage::I16(values) => select!(values, I16),
1254 IntegerStorage::I32(values) => select!(values, I32),
1255 IntegerStorage::I64(values) => select!(values, I64),
1256 IntegerStorage::U8(values) => select!(values, U8),
1257 IntegerStorage::U16(values) => select!(values, U16),
1258 IntegerStorage::U32(values) => select!(values, U32),
1259 IntegerStorage::U64(values) => select!(values, U64),
1260 }
1261}
1262
1263fn integer_values_to_value(
1264 storage: IntegerStorage,
1265 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
1266) -> crate::BuiltinResult<Value> {
1267 if storage.len() == 1 && output_provider.is_none() {
1268 return Ok(Value::Int(integer_storage_value(&storage, 0)));
1269 }
1270 let shape = vec![storage.len(), 1];
1271 let tensor = Tensor::new_integer(storage, shape)
1272 .map_err(|e| find_error_with_message(format!("find: {e}"), &FIND_ERROR_INTERNAL))?;
1273 Ok(tensor_to_value(tensor, output_provider))
1274}
1275
1276fn integer_storage_value(storage: &IntegerStorage, index: usize) -> IntValue {
1277 match storage {
1278 IntegerStorage::I8(values) => IntValue::I8(values[index]),
1279 IntegerStorage::I16(values) => IntValue::I16(values[index]),
1280 IntegerStorage::I32(values) => IntValue::I32(values[index]),
1281 IntegerStorage::I64(values) => IntValue::I64(values[index]),
1282 IntegerStorage::U8(values) => IntValue::U8(values[index]),
1283 IntegerStorage::U16(values) => IntValue::U16(values[index]),
1284 IntegerStorage::U32(values) => IntValue::U32(values[index]),
1285 IntegerStorage::U64(values) => IntValue::U64(values[index]),
1286 }
1287}
1288
1289fn integer_storage_from_scalar(value: &IntValue) -> IntegerStorage {
1290 match value {
1291 IntValue::I8(value) => IntegerStorage::I8(vec![*value]),
1292 IntValue::I16(value) => IntegerStorage::I16(vec![*value]),
1293 IntValue::I32(value) => IntegerStorage::I32(vec![*value]),
1294 IntValue::I64(value) => IntegerStorage::I64(vec![*value]),
1295 IntValue::U8(value) => IntegerStorage::U8(vec![*value]),
1296 IntValue::U16(value) => IntegerStorage::U16(vec![*value]),
1297 IntValue::U32(value) => IntegerStorage::U32(vec![*value]),
1298 IntValue::U64(value) => IntegerStorage::U64(vec![*value]),
1299 }
1300}
1301
1302fn tensor_to_value(
1303 tensor: Tensor,
1304 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
1305) -> Value {
1306 if let Some(provider) = output_provider {
1307 if let Ok(handle) = gpu_helpers::upload_tensor(provider, &tensor) {
1308 return Value::GpuTensor(handle);
1309 }
1310 }
1311 tensor::tensor_into_value(tensor)
1312}
1313
1314fn complex_tensor_to_value(
1315 tensor: ComplexTensor,
1316 output_provider: Option<&'static dyn runmat_accelerate_api::AccelProvider>,
1317) -> Value {
1318 if let Some(provider) = output_provider {
1319 if let Ok(handle) = gpu_helpers::upload_complex_tensor(provider, &tensor) {
1320 return gpu_helpers::complex_gpu_value(handle);
1321 }
1322 }
1323 complex_tensor_into_value(tensor)
1324}
1325
1326#[cfg(test)]
1327pub(crate) mod tests {
1328 use super::*;
1329 use crate::builtins::common::test_support;
1330 use futures::executor::block_on;
1331 use runmat_accelerate_api::HostTensorView;
1332 use runmat_builtins::Type;
1333 use runmat_value::{CharArray, IntValue};
1334
1335 fn find_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1336 block_on(super::find_builtin(value, rest))
1337 }
1338
1339 fn evaluate(value: Value, rest: &[Value]) -> crate::BuiltinResult<FindEval> {
1340 block_on(super::evaluate(value, rest))
1341 }
1342
1343 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1344 #[test]
1345 fn find_linear_indices_basic() {
1346 let tensor = Tensor::new(vec![0.0, 4.0, 0.0, 7.0, 0.0, 9.0], vec![2, 3]).unwrap();
1347 let value = find_builtin(Value::Tensor(tensor), Vec::new()).expect("find");
1348 match value {
1349 Value::Tensor(t) => {
1350 assert_eq!(t.shape, vec![3, 1]);
1351 assert_eq!(t.materialize_f64(), vec![2.0, 4.0, 6.0]);
1352 }
1353 other => panic!("expected tensor, got {other:?}"),
1354 }
1355 }
1356
1357 #[test]
1358 fn find_type_tracks_known_row_vector_orientation() {
1359 assert_eq!(
1360 find_type(
1361 &[Type::Tensor { shape: None }],
1362 &ResolveContext::new(Vec::new()),
1363 ),
1364 Type::Tensor {
1365 shape: Some(vec![None, Some(1)])
1366 }
1367 );
1368 assert_eq!(
1369 find_type(
1370 &[Type::Tensor {
1371 shape: Some(vec![Some(1), Some(5)])
1372 }],
1373 &ResolveContext::new(Vec::new()),
1374 ),
1375 Type::Tensor {
1376 shape: Some(vec![Some(1), None])
1377 }
1378 );
1379 }
1380
1381 #[test]
1382 fn find_integer_tokens_parse_exact_limits() {
1383 let options =
1384 parse_find_tokens(&[ArgToken::Integer(IntValue::U64(2))]).expect("uint64 limit");
1385 assert_eq!(options.limit, Some(2));
1386 assert_eq!(options.direction, FindDirection::First);
1387
1388 let options = parse_find_tokens(&[
1389 ArgToken::Integer(IntValue::U16(3)),
1390 ArgToken::String("last".to_string()),
1391 ])
1392 .expect("integer limit with direction");
1393 assert_eq!(options.limit, Some(3));
1394 assert_eq!(options.direction, FindDirection::Last);
1395
1396 let err = parse_find_tokens(&[ArgToken::Integer(IntValue::I64(-1))])
1397 .expect_err("negative integer limit must reject");
1398 assert_eq!(err.identifier(), FIND_ERROR_INVALID_INPUT.identifier);
1399 }
1400
1401 #[test]
1402 fn find_float_limits_reject_oversized_values_before_casting() {
1403 assert!(parse_find_tokens(&[ArgToken::Number(1.0e300)]).is_err());
1404 assert!(parse_find_tokens(&[ArgToken::Number(usize::MAX as f64)]).is_err());
1405 assert!(parse_find_tokens(&[ArgToken::Number(0.0)]).is_err());
1406 assert!(parse_find_tokens(&[ArgToken::Integer(IntValue::U8(0))]).is_err());
1407 }
1408
1409 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1410 #[test]
1411 fn find_limited_first() {
1412 let tensor = Tensor::new(vec![0.0, 3.0, 5.0, 0.0, 8.0], vec![1, 5]).unwrap();
1413 let result =
1414 find_builtin(Value::Tensor(tensor), vec![Value::Int(IntValue::I32(2))]).expect("find");
1415 match result {
1416 Value::Tensor(t) => {
1417 assert_eq!(t.shape, vec![1, 2]);
1418 assert_eq!(t.materialize_f64(), vec![2.0, 3.0]);
1419 }
1420 other => panic!("expected tensor, got {other:?}"),
1421 }
1422 }
1423
1424 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1425 #[test]
1426 fn find_last_single() {
1427 let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
1428 let tensor = Tensor::new(vec![1.0, 0.0, 0.0, 6.0, 0.0, 2.0], vec![1, 6]).unwrap();
1429 let result = find_builtin(Value::Tensor(tensor), vec![Value::from("last")]).expect("find");
1430 match result {
1431 Value::Num(n) => assert_eq!(n, 6.0),
1432 Value::Tensor(t) => {
1433 assert_eq!(t.materialize_f64(), vec![6.0]);
1434 }
1435 other => panic!("unexpected result {other:?}"),
1436 }
1437 }
1438
1439 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1440 #[test]
1441 fn find_complex_values() {
1442 let tensor =
1443 ComplexTensor::new(vec![(0.0, 0.0), (1.0, 2.0), (0.0, 0.0)], vec![3, 1]).unwrap();
1444 let eval = evaluate(Value::ComplexTensor(tensor), &[]).expect("find compute");
1445 let values = eval.values_value().expect("values");
1446 match values {
1447 Value::Complex(re, im) => {
1448 assert_eq!(re, 1.0);
1449 assert_eq!(im, 2.0);
1450 }
1451 Value::ComplexTensor(ct) => {
1452 assert_eq!(ct.shape, vec![1, 1]);
1453 assert_eq!(ct.materialize_f64(), vec![(1.0, 2.0)]);
1454 }
1455 other => panic!("expected complex result, got {other:?}"),
1456 }
1457 }
1458
1459 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1460 #[test]
1461 fn find_gpu_roundtrip() {
1462 test_support::with_test_provider(|provider| {
1463 let tensor = Tensor::new(vec![0.0, 4.0, 0.0, 7.0], vec![2, 2]).unwrap();
1464 let view = HostTensorView {
1465 data: &tensor.materialize_f64(),
1466 shape: &tensor.shape,
1467 };
1468 let handle = provider.upload(&view).expect("upload");
1469 let result = find_builtin(Value::GpuTensor(handle), Vec::new()).expect("find");
1470 let gathered = test_support::gather(result).expect("gather");
1471 assert_eq!(gathered.shape, vec![2, 1]);
1472 assert_eq!(gathered.materialize_f64(), vec![2.0, 4.0]);
1473 });
1474 }
1475
1476 #[test]
1477 fn find_f32_resident_fallback_returns_host_double_indices_when_owner_cannot_store_double() {
1478 test_support::with_f32_test_provider(|provider| {
1479 let values = [0.0, 4.0, 0.0, 7.0];
1480 let handle = provider
1481 .upload(&HostTensorView {
1482 data: &values,
1483 shape: &[2, 2],
1484 })
1485 .expect("upload f32-owner input");
1486
1487 let eval = evaluate(Value::GpuTensor(handle), &[]).expect("find fallback");
1488 let Value::Tensor(indices) = eval.linear_value().expect("linear indices") else {
1489 panic!("expected host double indices");
1490 };
1491 assert_eq!(indices.numeric_dtype(), runmat_value::NumericDType::F64);
1492 assert_eq!(indices.materialize_f64(), vec![2.0, 4.0]);
1493 });
1494 }
1495
1496 #[test]
1497 fn find_resident_logical_value_output_stays_logical_and_resident() {
1498 test_support::with_test_provider(|provider| {
1499 let values = [0.0, 1.0, 0.0, 1.0];
1500 let handle = provider
1501 .upload(&HostTensorView {
1502 data: &values,
1503 shape: &[2, 2],
1504 })
1505 .expect("upload logical input");
1506 let input = gpu_helpers::logical_gpu_value(handle);
1507
1508 let eval = evaluate(input, &[]).expect("find logical fallback");
1509 let Value::GpuTensor(values_handle) = eval.values_value().expect("selected values")
1510 else {
1511 panic!("expected resident logical selected values");
1512 };
1513 assert!(runmat_accelerate_api::handle_is_logical(&values_handle));
1514 let values =
1515 test_support::gather(Value::GpuTensor(values_handle)).expect("gather logical");
1516 assert_eq!(values.shape, vec![2, 1]);
1517 assert_eq!(values.materialize_f64(), vec![1.0, 1.0]);
1518 });
1519 }
1520
1521 #[test]
1522 fn find_routes_native_and_fallback_outputs_to_the_input_owner() {
1523 let _lock = test_support::accel_test_lock();
1524 let owner: &'static dyn runmat_accelerate_api::AccelProvider = Box::leak(Box::new(
1525 runmat_accelerate::simple_provider::InProcessProvider::new(),
1526 ));
1527 let active: &'static dyn runmat_accelerate_api::AccelProvider = Box::leak(Box::new(
1528 runmat_accelerate::simple_provider::InProcessProvider::new(),
1529 ));
1530 unsafe {
1531 runmat_accelerate_api::register_provider(owner);
1532 runmat_accelerate_api::register_provider(active);
1533 }
1534 let _active = runmat_accelerate_api::ThreadProviderGuard::set(Some(active));
1535 assert_ne!(owner.device_id(), active.device_id());
1536
1537 let native_input = owner
1538 .upload(&HostTensorView {
1539 data: &[0.0, 4.0, 0.0, 7.0],
1540 shape: &[2, 2],
1541 })
1542 .expect("upload native input");
1543 let native = evaluate(Value::GpuTensor(native_input), &[]).expect("native find");
1544 let Value::GpuTensor(native_indices) = native.linear_value().expect("native indices")
1545 else {
1546 panic!("expected native resident indices");
1547 };
1548 assert_eq!(native_indices.device_id, owner.device_id());
1549
1550 let fallback_input = owner
1551 .upload(&HostTensorView {
1552 data: &[0.0, 4.0, 0.0, 7.0],
1553 shape: &[2, 2],
1554 })
1555 .expect("upload fallback input");
1556 let fallback = evaluate(Value::GpuTensor(fallback_input), &[]).expect("fallback find");
1557 let Value::GpuTensor(fallback_indices) = fallback.linear_value().expect("fallback indices")
1558 else {
1559 panic!("expected fallback resident indices");
1560 };
1561 assert_eq!(fallback_indices.device_id, owner.device_id());
1562 assert_eq!(
1563 test_support::gather(Value::GpuTensor(fallback_indices))
1564 .expect("gather fallback indices")
1565 .materialize_f64(),
1566 vec![2.0, 4.0]
1567 );
1568
1569 let integer_input = owner
1570 .upload_integer(&runmat_accelerate_api::HostIntegerTensorView {
1571 data: runmat_accelerate_api::HostIntegerDataView::U64(&[0, 9_007_199_254_740_993]),
1572 shape: &[1, 2],
1573 })
1574 .expect("upload integer input");
1575 let integer = evaluate(Value::GpuTensor(integer_input), &[]).expect("integer find");
1576 let Value::GpuTensor(integer_values) = integer.values_value().expect("integer values")
1577 else {
1578 panic!("expected resident integer values");
1579 };
1580 assert_eq!(integer_values.device_id, owner.device_id());
1581 assert_eq!(
1582 runmat_accelerate_api::handle_integer_type(&integer_values),
1583 Some(runmat_accelerate_api::IntegerElementType::U64)
1584 );
1585 assert_eq!(
1586 test_support::gather(Value::GpuTensor(integer_values))
1587 .expect("gather integer values")
1588 .integer_storage(),
1589 Some(&IntegerStorage::U64(vec![9_007_199_254_740_993]))
1590 );
1591
1592 let logical_input = owner
1593 .upload(&HostTensorView {
1594 data: &[0.0, 1.0],
1595 shape: &[1, 2],
1596 })
1597 .expect("upload logical input");
1598 let logical =
1599 evaluate(gpu_helpers::logical_gpu_value(logical_input), &[]).expect("logical find");
1600 let Value::GpuTensor(logical_values) = logical.values_value().expect("logical values")
1601 else {
1602 panic!("expected resident logical values");
1603 };
1604 assert_eq!(logical_values.device_id, owner.device_id());
1605 assert!(runmat_accelerate_api::handle_is_logical(&logical_values));
1606 }
1607
1608 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1609 #[test]
1610 fn find_gpu_row_vector_preserves_linear_index_orientation() {
1611 test_support::with_test_provider(|provider| {
1612 let tensor = Tensor::new(vec![0.0, 4.0, 5.0, 7.0], vec![1, 4]).unwrap();
1613 let view = HostTensorView {
1614 data: &tensor.materialize_f64(),
1615 shape: &tensor.shape,
1616 };
1617 let handle = provider.upload(&view).expect("upload");
1618 let result = find_builtin(Value::GpuTensor(handle), Vec::new()).expect("find");
1619 assert!(matches!(result, Value::GpuTensor(_)));
1620 let gathered = test_support::gather(result).expect("gather");
1621 assert_eq!(gathered.shape, vec![1, 3]);
1622 assert_eq!(gathered.materialize_f64(), vec![2.0, 3.0, 4.0]);
1623 });
1624 }
1625
1626 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1627 #[test]
1628 fn find_direction_error() {
1629 let tensor = Tensor::new(vec![1.0], vec![1, 1]).unwrap();
1630 let err = find_builtin(
1631 Value::Tensor(tensor),
1632 vec![Value::Int(IntValue::I32(1)), Value::from("invalid")],
1633 )
1634 .expect_err("expected error");
1635 assert!(err.to_string().contains("direction"));
1636 assert_eq!(err.identifier(), super::FIND_ERROR_INVALID_INPUT.identifier);
1637 }
1638
1639 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1640 #[test]
1641 fn find_multi_output_rows_cols_values() {
1642 let tensor = Tensor::new(vec![0.0, 2.0, 3.0, 0.0, 0.0, 6.0], vec![2, 3]).unwrap();
1643 let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1644
1645 let rows = test_support::gather(eval.row_value().expect("rows")).expect("gather rows");
1646 assert_eq!(rows.shape, vec![3, 1]);
1647 assert_eq!(rows.materialize_f64(), vec![2.0, 1.0, 2.0]);
1648
1649 let cols = test_support::gather(eval.column_value().expect("cols")).expect("gather cols");
1650 assert_eq!(cols.shape, vec![3, 1]);
1651 assert_eq!(cols.materialize_f64(), vec![1.0, 2.0, 3.0]);
1652
1653 let vals = test_support::gather(eval.values_value().expect("vals")).expect("gather vals");
1654 assert_eq!(vals.shape, vec![3, 1]);
1655 assert_eq!(vals.materialize_f64(), vec![2.0, 3.0, 6.0]);
1656 }
1657
1658 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1659 #[test]
1660 fn find_values_preserve_exact_uint64_storage() {
1661 let input = Tensor::new_integer(
1662 IntegerStorage::U64(vec![0, u64::MAX, 1_u64 << 63, 0]),
1663 vec![2, 2],
1664 )
1665 .expect("integer tensor");
1666 let eval = evaluate(Value::Tensor(input), &[]).expect("evaluate");
1667 let values = eval.values_value().expect("values");
1668 let Value::Tensor(values) = values else {
1669 panic!("expected typed tensor values");
1670 };
1671 assert_eq!(values.shape, vec![2, 1]);
1672 assert_eq!(
1673 values.integer_storage(),
1674 Some(&IntegerStorage::U64(vec![u64::MAX, 1_u64 << 63]))
1675 );
1676 }
1677
1678 #[test]
1679 fn find_indices_read_typed_integer_storage_exactly() {
1680 let input = Tensor::new_integer(IntegerStorage::I16(vec![0, -7, 0, 9]), vec![2, 2])
1681 .expect("integer tensor");
1682
1683 let value = find_builtin(Value::Tensor(input), Vec::new()).expect("find");
1684
1685 match value {
1686 Value::Tensor(indices) => {
1687 assert_eq!(indices.shape, vec![2, 1]);
1688 assert_eq!(indices.materialize_f64(), vec![2.0, 4.0]);
1689 }
1690 other => panic!("expected index tensor, got {other:?}"),
1691 }
1692 }
1693
1694 #[test]
1695 fn find_last_indices_read_typed_integer_storage_exactly() {
1696 let input = Tensor::new_integer(IntegerStorage::U16(vec![5, 0, 3, 0]), vec![2, 2])
1697 .expect("integer tensor");
1698
1699 let value = find_builtin(
1700 Value::Tensor(input),
1701 vec![Value::Int(IntValue::I32(1)), Value::from("last")],
1702 )
1703 .expect("find");
1704
1705 assert_eq!(value, Value::Num(3.0));
1706 }
1707
1708 #[test]
1709 fn find_reads_mirrorless_typed_complex_integer_storage() {
1710 let storage = IntegerComplexStorage::new(
1711 IntegerStorage::I16(vec![0, -7, 0, 9]),
1712 IntegerStorage::I16(vec![0, 0, 5, 0]),
1713 )
1714 .expect("complex integer storage");
1715 let input = ComplexTensor::new_integer(storage, vec![2, 2]).expect("complex tensor");
1716
1717 let eval = evaluate(Value::ComplexTensor(input), &[]).expect("find");
1718 let linear = tensor::value_into_tensor_for("find", eval.linear_value().expect("linear"))
1719 .expect("linear tensor");
1720 assert_eq!(linear.materialize_f64(), vec![2.0, 3.0, 4.0]);
1721 let values = eval.values_value().expect("values");
1722 let Value::ComplexTensor(values) = values else {
1723 panic!("expected typed complex tensor values");
1724 };
1725 let storage = values.integer_storage().expect("typed complex values");
1726 assert_eq!(storage.real, IntegerStorage::I16(vec![-7, 0, 9]));
1727 assert_eq!(storage.imag, IntegerStorage::I16(vec![0, 5, 0]));
1728 }
1729
1730 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1731 #[test]
1732 fn find_sparse_values_preserve_exact_storage_and_traversal_order() {
1733 let _extensions = crate::compatibility::push_runmat_extensions_enabled(true);
1734 let sparse = runmat_value::SparseTensor::new_integer(
1735 3,
1736 2,
1737 vec![0, 2, 3],
1738 vec![0, 2, 1],
1739 IntegerStorage::U64(vec![u64::MAX, 1_u64 << 63, 7]),
1740 )
1741 .expect("typed sparse");
1742
1743 let all = evaluate(Value::SparseTensor(sparse.clone()), &[]).expect("find sparse");
1744 let Value::Tensor(all_values) = all.values_value().expect("all values") else {
1745 panic!("expected typed sparse find values");
1746 };
1747 assert_eq!(
1748 all_values.integer_storage(),
1749 Some(&IntegerStorage::U64(vec![u64::MAX, 1_u64 << 63, 7]))
1750 );
1751
1752 let first = evaluate(
1753 Value::SparseTensor(sparse.clone()),
1754 &[Value::Int(IntValue::I32(2))],
1755 )
1756 .expect("find sparse first");
1757 let Value::Tensor(first_values) = first.values_value().expect("first values") else {
1758 panic!("expected typed sparse first values");
1759 };
1760 assert_eq!(
1761 first_values.integer_storage(),
1762 Some(&IntegerStorage::U64(vec![u64::MAX, 1_u64 << 63]))
1763 );
1764
1765 let last = evaluate(
1766 Value::SparseTensor(sparse),
1767 &[Value::Int(IntValue::I32(2)), Value::from("last")],
1768 )
1769 .expect("find sparse last");
1770 let Value::Tensor(last_values) = last.values_value().expect("last values") else {
1771 panic!("expected typed sparse last values");
1772 };
1773 assert_eq!(
1774 last_values.integer_storage(),
1775 Some(&IntegerStorage::U64(vec![1_u64 << 63, 7]))
1776 );
1777 }
1778
1779 #[test]
1780 fn find_sparse_values_preserve_native_single_class_and_order() {
1781 let sparse = runmat_value::SparseTensor::new_f32(
1782 3,
1783 2,
1784 vec![0, 2, 3],
1785 vec![0, 2, 1],
1786 vec![1.25, 3.5, 7.75],
1787 )
1788 .expect("single sparse");
1789 let eval = evaluate(Value::SparseTensor(sparse), &[]).expect("find sparse");
1790 let Value::Tensor(values) = eval.values_value().expect("values") else {
1791 panic!("expected native-single find values");
1792 };
1793 assert_eq!(values.numeric_dtype(), runmat_value::NumericDType::F32);
1794 assert_eq!(values.as_f32_slice(), Some(&[1.25, 3.5, 7.75][..]));
1795 }
1796
1797 #[test]
1798 fn find_sparse_values_preserve_logical_class_and_order() {
1799 let sparse = runmat_value::SparseTensor::new_logical(3, 2, vec![0, 2, 3], vec![0, 2, 1])
1800 .expect("logical sparse");
1801 let eval = evaluate(Value::SparseTensor(sparse), &[]).expect("find sparse");
1802 let Value::LogicalArray(values) = eval.values_value().expect("values") else {
1803 panic!("expected logical sparse find values");
1804 };
1805 assert_eq!(values.shape, vec![3, 1]);
1806 assert_eq!(values.data, vec![1, 1, 1]);
1807 }
1808
1809 #[test]
1810 fn find_integer_selection_preserves_every_integer_class() {
1811 let cases = [
1812 (
1813 IntegerStorage::I8(vec![-2, 0, 3]),
1814 IntegerStorage::I8(vec![3, -2]),
1815 ),
1816 (
1817 IntegerStorage::I16(vec![-2, 0, 3]),
1818 IntegerStorage::I16(vec![3, -2]),
1819 ),
1820 (
1821 IntegerStorage::I32(vec![-2, 0, 3]),
1822 IntegerStorage::I32(vec![3, -2]),
1823 ),
1824 (
1825 IntegerStorage::I64(vec![-2, 0, 3]),
1826 IntegerStorage::I64(vec![3, -2]),
1827 ),
1828 (
1829 IntegerStorage::U8(vec![2, 0, 3]),
1830 IntegerStorage::U8(vec![3, 2]),
1831 ),
1832 (
1833 IntegerStorage::U16(vec![2, 0, 3]),
1834 IntegerStorage::U16(vec![3, 2]),
1835 ),
1836 (
1837 IntegerStorage::U32(vec![2, 0, 3]),
1838 IntegerStorage::U32(vec![3, 2]),
1839 ),
1840 (
1841 IntegerStorage::U64(vec![2, 0, 3]),
1842 IntegerStorage::U64(vec![3, 2]),
1843 ),
1844 ];
1845 for (storage, expected) in cases {
1846 assert_eq!(select_integer_values(&storage, &[2, 0]), expected);
1847 }
1848 }
1849
1850 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1851 #[test]
1852 fn find_single_integer_value_preserves_scalar_class() {
1853 let input = Tensor::new_integer(IntegerStorage::I64(vec![0, i64::MIN]), vec![2, 1])
1854 .expect("integer tensor");
1855 let eval = evaluate(Value::Tensor(input), &[]).expect("evaluate");
1856 assert_eq!(
1857 eval.values_value().expect("values"),
1858 Value::Int(IntValue::I64(i64::MIN))
1859 );
1860 }
1861
1862 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1863 #[test]
1864 fn find_integer_scalar_preserves_exact_value_output() {
1865 let eval = evaluate(Value::Int(IntValue::U64(u64::MAX)), &[]).expect("evaluate");
1866 assert_eq!(
1867 eval.values_value().expect("values"),
1868 Value::Int(IntValue::U64(u64::MAX))
1869 );
1870 }
1871
1872 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1873 #[test]
1874 fn find_last_returns_selected_indices_in_ascending_order() {
1875 let tensor = Tensor::new(vec![1.0, 0.0, 2.0, 3.0, 0.0], vec![1, 5]).unwrap();
1876 let result = find_builtin(
1877 Value::Tensor(tensor),
1878 vec![Value::Int(IntValue::I32(2)), Value::from("last")],
1879 )
1880 .expect("find");
1881 match result {
1882 Value::Tensor(t) => {
1883 assert_eq!(t.shape, vec![1, 2]);
1884 assert_eq!(t.materialize_f64(), vec![3.0, 4.0]);
1885 }
1886 Value::Num(_) => panic!("expected column vector"),
1887 other => panic!("unexpected result {other:?}"),
1888 }
1889 }
1890
1891 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1892 #[test]
1893 fn find_limit_zero_rejects() {
1894 let tensor = Tensor::new(vec![1.0, 0.0, 3.0], vec![3, 1]).unwrap();
1895 find_builtin(Value::Tensor(tensor), vec![Value::Num(0.0)])
1896 .expect_err("zero is not a positive count");
1897 }
1898
1899 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1900 #[test]
1901 fn find_empty_orientation_follows_input_vector_shape() {
1902 for (shape, expected_shape) in [
1903 (vec![1, 0], vec![1, 0]),
1904 (vec![0, 1], vec![0, 1]),
1905 (vec![0, 3], vec![0, 1]),
1906 ] {
1907 let input =
1908 Tensor::new_integer(IntegerStorage::U32(Vec::new()), shape).expect("empty input");
1909 let Value::Tensor(indices) =
1910 find_builtin(Value::Tensor(input), Vec::new()).expect("find")
1911 else {
1912 panic!("expected empty tensor");
1913 };
1914 assert_eq!(indices.shape, expected_shape);
1915 assert!(indices.materialize_f64().is_empty());
1916 }
1917 }
1918
1919 #[test]
1920 fn find_scalar_zero_and_empty_matrix_use_empty_matrix_convention() {
1921 for input in [
1922 Value::Num(0.0),
1923 Value::Tensor(Tensor::new(Vec::new(), vec![0, 0]).unwrap()),
1924 ] {
1925 let Value::Tensor(indices) = find_builtin(input, Vec::new()).expect("find") else {
1926 panic!("expected empty tensor");
1927 };
1928 assert_eq!(indices.shape, vec![0, 0]);
1929 }
1930 }
1931
1932 #[test]
1933 fn find_dense_logical_value_output_preserves_logical_class() {
1934 let input = LogicalArray::new(vec![0, 1, 1, 0], vec![2, 2]).unwrap();
1935 let eval = evaluate(Value::LogicalArray(input), &[]).expect("find");
1936 let Value::LogicalArray(values) = eval.values_value().expect("values") else {
1937 panic!("expected logical values");
1938 };
1939 assert_eq!(values.shape, vec![2, 1]);
1940 assert_eq!(values.data, vec![1, 1]);
1941 }
1942
1943 #[test]
1944 fn find_runmat_only_forms_gate_before_evaluation() {
1945 let _strict = crate::compatibility::push_runmat_extensions_enabled(false);
1946 let input = Value::Tensor(Tensor::new(vec![0.0, 1.0], vec![1, 2]).unwrap());
1947 let err = evaluate(input, &[Value::from("last")])
1948 .err()
1949 .expect("direction-only form must gate");
1950 assert_eq!(
1951 err.identifier(),
1952 FIND_DIRECTION_ONLY_EXTENSION.error_identifier
1953 );
1954
1955 let sparse = runmat_value::SparseTensor::new_integer(
1956 1,
1957 1,
1958 vec![0, 1],
1959 vec![0],
1960 IntegerStorage::U64(vec![u64::MAX]),
1961 )
1962 .unwrap();
1963 let err = evaluate(Value::SparseTensor(sparse), &[])
1964 .err()
1965 .expect("integer sparse form must gate");
1966 assert_eq!(
1967 err.identifier(),
1968 FIND_INTEGER_SPARSE_EXTENSION.error_identifier
1969 );
1970 }
1971
1972 #[test]
1973 fn find_integer_metadata_covers_values_counts_and_sparse_extension() {
1974 assert_eq!(FIND_INTEGER_CAPABILITIES.len(), 4);
1975 assert_eq!(FIND_EXTENSIONS.len(), 2);
1976 for capability in FIND_INTEGER_CAPABILITIES {
1977 for input in capability.inputs {
1978 assert_eq!(
1979 input.classes,
1980 &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES
1981 );
1982 }
1983 }
1984 if let Some(largest_exact_index) = 1_usize.checked_shl(53) {
1985 assert_eq!(
1986 exact_index_as_f64(largest_exact_index).expect("largest exact index"),
1987 9_007_199_254_740_992.0
1988 );
1989 }
1990 }
1991
1992 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1993 #[test]
1994 fn find_integer_gpu_last_preserves_order_orientation_class_and_residency() {
1995 test_support::with_f32_test_provider(|provider| {
1996 let handle = provider
1997 .upload_integer(&runmat_accelerate_api::HostIntegerTensorView {
1998 data: runmat_accelerate_api::HostIntegerDataView::U64(&[
1999 0,
2000 1_u64 << 63,
2001 7,
2002 u64::MAX,
2003 ]),
2004 shape: &[1, 4],
2005 })
2006 .expect("upload integer row vector");
2007 let eval = evaluate(
2008 Value::GpuTensor(handle),
2009 &[Value::Int(IntValue::I32(2)), Value::from("last")],
2010 )
2011 .expect("find last");
2012
2013 let linear = eval.linear_value().expect("linear indices");
2014 let Value::Tensor(linear) = linear else {
2015 panic!("double indices must fall back to host storage");
2016 };
2017 assert_eq!(linear.shape, vec![1, 2]);
2018 assert_eq!(linear.materialize_f64(), vec![3.0, 4.0]);
2019
2020 let values = eval.values_value().expect("selected values");
2021 let Value::GpuTensor(values_handle) = &values else {
2022 panic!("expected resident selected values, got {values:?}");
2023 };
2024 assert_eq!(
2025 runmat_accelerate_api::handle_integer_type(values_handle),
2026 Some(runmat_accelerate_api::IntegerElementType::U64)
2027 );
2028 let values = test_support::gather(values).expect("gather selected values");
2029 assert_eq!(values.shape, vec![2, 1]);
2030 assert_eq!(
2031 values.integer_storage(),
2032 Some(&IntegerStorage::U64(vec![7, u64::MAX]))
2033 );
2034 });
2035 }
2036
2037 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2038 #[test]
2039 fn find_char_array_supports_nonzero_codes() {
2040 let chars = CharArray::new(vec!['\0', 'A', '\0'], 1, 3).unwrap();
2041 let result = find_builtin(Value::CharArray(chars), Vec::new()).expect("find");
2042 match result {
2043 Value::Num(n) => assert_eq!(n, 2.0),
2044 Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0]),
2045 other => panic!("unexpected result {other:?}"),
2046 }
2047 }
2048
2049 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2050 #[test]
2051 fn find_gpu_multi_outputs_return_gpu_handles() {
2052 test_support::with_test_provider(|provider| {
2053 let tensor = Tensor::new(vec![0.0, 4.0, 5.0, 0.0], vec![2, 2]).unwrap();
2054 let view = HostTensorView {
2055 data: &tensor.materialize_f64(),
2056 shape: &tensor.shape,
2057 };
2058 let handle = provider.upload(&view).expect("upload");
2059 let eval = evaluate(Value::GpuTensor(handle), &[]).expect("evaluate");
2060
2061 let rows = eval.row_value().expect("rows");
2062 assert!(matches!(rows, Value::GpuTensor(_)));
2063 let rows_host = test_support::gather(rows).expect("gather rows");
2064 assert_eq!(rows_host.materialize_f64(), vec![2.0, 1.0]);
2065
2066 let cols = eval.column_value().expect("cols");
2067 assert!(matches!(cols, Value::GpuTensor(_)));
2068 let cols_host = test_support::gather(cols).expect("gather cols");
2069 assert_eq!(cols_host.materialize_f64(), vec![1.0, 2.0]);
2070
2071 let vals = eval.values_value().expect("vals");
2072 assert!(matches!(vals, Value::GpuTensor(_)));
2073 let vals_host = test_support::gather(vals).expect("gather vals");
2074 assert_eq!(vals_host.materialize_f64(), vec![4.0, 5.0]);
2075 });
2076 }
2077
2078 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2079 #[test]
2080 #[cfg(feature = "wgpu")]
2081 fn find_wgpu_matches_cpu() {
2082 let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2083 runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2084 );
2085 let tensor = Tensor::new(vec![0.0, 2.0, 0.0, 3.0, 4.0, 0.0], vec![3, 2]).unwrap();
2086 let cpu_eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("cpu evaluate");
2087 let cpu_linear =
2088 test_support::gather(cpu_eval.linear_value().expect("cpu linear")).expect("cpu gather");
2089 let provider = runmat_accelerate_api::provider().expect("wgpu provider");
2090 let view = HostTensorView {
2091 data: &tensor.materialize_f64(),
2092 shape: &tensor.shape,
2093 };
2094 let handle = provider.upload(&view).expect("upload");
2095 let gpu_eval = evaluate(Value::GpuTensor(handle), &[]).expect("gpu evaluate");
2096 let gpu_linear =
2097 test_support::gather(gpu_eval.linear_value().expect("gpu linear")).expect("gpu gather");
2098 assert_eq!(gpu_linear.materialize_f64(), cpu_linear.materialize_f64());
2099 }
2100}