Skip to main content

runmat_runtime/builtins/array/indexing/
find.rs

1//! MATLAB-compatible `find` builtin with GPU-aware semantics for RunMat.
2
3use 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
440/// Evaluate `find` and return an object that can materialise the various outputs.
441pub 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}