Skip to main content

runmat_runtime/builtins/math/interpolation/
interp2.rs

1//! MATLAB-compatible `interp2` builtin for gridded dense real data.
2
3use runmat_builtins::{
4    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
5    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
6    ResolveContext, Tensor, Type, Value,
7};
8use runmat_macros::runtime_builtin;
9
10use crate::builtins::common::spec::{
11    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
12    ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
13};
14use crate::builtins::common::tensor;
15use crate::dispatcher;
16use crate::{build_runtime_error, RuntimeError};
17
18use super::pp::{
19    interval_index, is_vector_shape, out_of_range_value, parse_extrapolation, parse_method,
20    query_points, vector_from_value, Extrapolation, InterpMethod,
21};
22
23const NAME: &str = "interp2";
24
25const INTERP2_OUTPUT: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
26    name: "Vq",
27    ty: BuiltinParamType::NumericArray,
28    arity: BuiltinParamArity::Required,
29    default: None,
30    description: "Interpolated values over a 2-D grid.",
31}];
32
33const INTERP2_INPUTS_Z_XQ_YQ: [BuiltinParamDescriptor; 3] = [
34    BuiltinParamDescriptor {
35        name: "Z",
36        ty: BuiltinParamType::Any,
37        arity: BuiltinParamArity::Required,
38        default: None,
39        description: "Grid sample matrix.",
40    },
41    BuiltinParamDescriptor {
42        name: "Xq",
43        ty: BuiltinParamType::Any,
44        arity: BuiltinParamArity::Required,
45        default: None,
46        description: "Query X coordinates.",
47    },
48    BuiltinParamDescriptor {
49        name: "Yq",
50        ty: BuiltinParamType::Any,
51        arity: BuiltinParamArity::Required,
52        default: None,
53        description: "Query Y coordinates.",
54    },
55];
56
57const INTERP2_INPUTS_X_Y_Z_XQ_YQ: [BuiltinParamDescriptor; 5] = [
58    BuiltinParamDescriptor {
59        name: "X",
60        ty: BuiltinParamType::Any,
61        arity: BuiltinParamArity::Required,
62        default: None,
63        description: "Grid X axis vector or mesh.",
64    },
65    BuiltinParamDescriptor {
66        name: "Y",
67        ty: BuiltinParamType::Any,
68        arity: BuiltinParamArity::Required,
69        default: None,
70        description: "Grid Y axis vector or mesh.",
71    },
72    BuiltinParamDescriptor {
73        name: "Z",
74        ty: BuiltinParamType::Any,
75        arity: BuiltinParamArity::Required,
76        default: None,
77        description: "Grid sample matrix.",
78    },
79    BuiltinParamDescriptor {
80        name: "Xq",
81        ty: BuiltinParamType::Any,
82        arity: BuiltinParamArity::Required,
83        default: None,
84        description: "Query X coordinates.",
85    },
86    BuiltinParamDescriptor {
87        name: "Yq",
88        ty: BuiltinParamType::Any,
89        arity: BuiltinParamArity::Required,
90        default: None,
91        description: "Query Y coordinates.",
92    },
93];
94
95const INTERP2_INPUTS_Z_XQ_YQ_METHOD: [BuiltinParamDescriptor; 4] = [
96    BuiltinParamDescriptor {
97        name: "Z",
98        ty: BuiltinParamType::Any,
99        arity: BuiltinParamArity::Required,
100        default: None,
101        description: "Grid sample matrix.",
102    },
103    BuiltinParamDescriptor {
104        name: "Xq",
105        ty: BuiltinParamType::Any,
106        arity: BuiltinParamArity::Required,
107        default: None,
108        description: "Query X coordinates.",
109    },
110    BuiltinParamDescriptor {
111        name: "Yq",
112        ty: BuiltinParamType::Any,
113        arity: BuiltinParamArity::Required,
114        default: None,
115        description: "Query Y coordinates.",
116    },
117    BuiltinParamDescriptor {
118        name: "method",
119        ty: BuiltinParamType::StringScalar,
120        arity: BuiltinParamArity::Optional,
121        default: Some("\"linear\""),
122        description: "Interpolation method: \"linear\" or \"nearest\".",
123    },
124];
125
126const INTERP2_INPUTS_Z_XQ_YQ_METHOD_EXTRAP: [BuiltinParamDescriptor; 5] = [
127    BuiltinParamDescriptor {
128        name: "Z",
129        ty: BuiltinParamType::Any,
130        arity: BuiltinParamArity::Required,
131        default: None,
132        description: "Grid sample matrix.",
133    },
134    BuiltinParamDescriptor {
135        name: "Xq",
136        ty: BuiltinParamType::Any,
137        arity: BuiltinParamArity::Required,
138        default: None,
139        description: "Query X coordinates.",
140    },
141    BuiltinParamDescriptor {
142        name: "Yq",
143        ty: BuiltinParamType::Any,
144        arity: BuiltinParamArity::Required,
145        default: None,
146        description: "Query Y coordinates.",
147    },
148    BuiltinParamDescriptor {
149        name: "method",
150        ty: BuiltinParamType::StringScalar,
151        arity: BuiltinParamArity::Optional,
152        default: Some("\"linear\""),
153        description: "Interpolation method: \"linear\" or \"nearest\".",
154    },
155    BuiltinParamDescriptor {
156        name: "extrap",
157        ty: BuiltinParamType::Any,
158        arity: BuiltinParamArity::Optional,
159        default: Some("NaN"),
160        description: "Extrapolation mode: \"extrap\" or scalar fill value.",
161    },
162];
163
164const INTERP2_INPUTS_X_Y_Z_XQ_YQ_METHOD: [BuiltinParamDescriptor; 6] = [
165    BuiltinParamDescriptor {
166        name: "X",
167        ty: BuiltinParamType::Any,
168        arity: BuiltinParamArity::Required,
169        default: None,
170        description: "Grid X axis vector or mesh.",
171    },
172    BuiltinParamDescriptor {
173        name: "Y",
174        ty: BuiltinParamType::Any,
175        arity: BuiltinParamArity::Required,
176        default: None,
177        description: "Grid Y axis vector or mesh.",
178    },
179    BuiltinParamDescriptor {
180        name: "Z",
181        ty: BuiltinParamType::Any,
182        arity: BuiltinParamArity::Required,
183        default: None,
184        description: "Grid sample matrix.",
185    },
186    BuiltinParamDescriptor {
187        name: "Xq",
188        ty: BuiltinParamType::Any,
189        arity: BuiltinParamArity::Required,
190        default: None,
191        description: "Query X coordinates.",
192    },
193    BuiltinParamDescriptor {
194        name: "Yq",
195        ty: BuiltinParamType::Any,
196        arity: BuiltinParamArity::Required,
197        default: None,
198        description: "Query Y coordinates.",
199    },
200    BuiltinParamDescriptor {
201        name: "method",
202        ty: BuiltinParamType::StringScalar,
203        arity: BuiltinParamArity::Optional,
204        default: Some("\"linear\""),
205        description: "Interpolation method: \"linear\" or \"nearest\".",
206    },
207];
208
209const INTERP2_INPUTS_X_Y_Z_XQ_YQ_METHOD_EXTRAP: [BuiltinParamDescriptor; 7] = [
210    BuiltinParamDescriptor {
211        name: "X",
212        ty: BuiltinParamType::Any,
213        arity: BuiltinParamArity::Required,
214        default: None,
215        description: "Grid X axis vector or mesh.",
216    },
217    BuiltinParamDescriptor {
218        name: "Y",
219        ty: BuiltinParamType::Any,
220        arity: BuiltinParamArity::Required,
221        default: None,
222        description: "Grid Y axis vector or mesh.",
223    },
224    BuiltinParamDescriptor {
225        name: "Z",
226        ty: BuiltinParamType::Any,
227        arity: BuiltinParamArity::Required,
228        default: None,
229        description: "Grid sample matrix.",
230    },
231    BuiltinParamDescriptor {
232        name: "Xq",
233        ty: BuiltinParamType::Any,
234        arity: BuiltinParamArity::Required,
235        default: None,
236        description: "Query X coordinates.",
237    },
238    BuiltinParamDescriptor {
239        name: "Yq",
240        ty: BuiltinParamType::Any,
241        arity: BuiltinParamArity::Required,
242        default: None,
243        description: "Query Y coordinates.",
244    },
245    BuiltinParamDescriptor {
246        name: "method",
247        ty: BuiltinParamType::StringScalar,
248        arity: BuiltinParamArity::Optional,
249        default: Some("\"linear\""),
250        description: "Interpolation method: \"linear\" or \"nearest\".",
251    },
252    BuiltinParamDescriptor {
253        name: "extrap",
254        ty: BuiltinParamType::Any,
255        arity: BuiltinParamArity::Optional,
256        default: Some("NaN"),
257        description: "Extrapolation mode: \"extrap\" or scalar fill value.",
258    },
259];
260
261const INTERP2_SIGNATURES: [BuiltinSignatureDescriptor; 8] = [
262    BuiltinSignatureDescriptor {
263        label: "Vq = interp2(Z, Xq, Yq)",
264        inputs: &INTERP2_INPUTS_Z_XQ_YQ,
265        outputs: &INTERP2_OUTPUT,
266    },
267    BuiltinSignatureDescriptor {
268        label: "Vq = interp2(X, Y, Z, Xq, Yq)",
269        inputs: &INTERP2_INPUTS_X_Y_Z_XQ_YQ,
270        outputs: &INTERP2_OUTPUT,
271    },
272    BuiltinSignatureDescriptor {
273        label: "Vq = interp2(Z, Xq, Yq, method)",
274        inputs: &INTERP2_INPUTS_Z_XQ_YQ_METHOD,
275        outputs: &INTERP2_OUTPUT,
276    },
277    BuiltinSignatureDescriptor {
278        label: "Vq = interp2(X, Y, Z, Xq, Yq, method)",
279        inputs: &INTERP2_INPUTS_X_Y_Z_XQ_YQ_METHOD,
280        outputs: &INTERP2_OUTPUT,
281    },
282    BuiltinSignatureDescriptor {
283        label: "Vq = interp2(Z, Xq, Yq, extrap)",
284        inputs: &INTERP2_INPUTS_Z_XQ_YQ_METHOD,
285        outputs: &INTERP2_OUTPUT,
286    },
287    BuiltinSignatureDescriptor {
288        label: "Vq = interp2(X, Y, Z, Xq, Yq, extrap)",
289        inputs: &INTERP2_INPUTS_X_Y_Z_XQ_YQ_METHOD,
290        outputs: &INTERP2_OUTPUT,
291    },
292    BuiltinSignatureDescriptor {
293        label: "Vq = interp2(Z, Xq, Yq, method, extrap)",
294        inputs: &INTERP2_INPUTS_Z_XQ_YQ_METHOD_EXTRAP,
295        outputs: &INTERP2_OUTPUT,
296    },
297    BuiltinSignatureDescriptor {
298        label: "Vq = interp2(X, Y, Z, Xq, Yq, method, extrap)",
299        inputs: &INTERP2_INPUTS_X_Y_Z_XQ_YQ_METHOD_EXTRAP,
300        outputs: &INTERP2_OUTPUT,
301    },
302];
303
304const INTERP2_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
305    code: "RM.INTERP2.INVALID_ARGUMENT",
306    identifier: Some("RunMat:interp2:InvalidArgument"),
307    when: "Argument count, method/extrapolation options, or axis/query compatibility is invalid.",
308    message: "interp2: invalid argument",
309};
310
311const INTERP2_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
312    code: "RM.INTERP2.INVALID_INPUT",
313    identifier: Some("RunMat:interp2:InvalidInput"),
314    when: "Grid or query values cannot be converted to numeric interpolation domains.",
315    message: "interp2: invalid input",
316};
317
318const INTERP2_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
319    code: "RM.INTERP2.INTERNAL",
320    identifier: Some("RunMat:interp2:Internal"),
321    when: "Interpolation output construction fails due to internal tensor assembly paths.",
322    message: "interp2: internal interpolation failure",
323};
324
325const INTERP2_ERRORS: [BuiltinErrorDescriptor; 3] = [
326    INTERP2_ERROR_INVALID_ARGUMENT,
327    INTERP2_ERROR_INVALID_INPUT,
328    INTERP2_ERROR_INTERNAL,
329];
330
331pub const INTERP2_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
332    signatures: &INTERP2_SIGNATURES,
333    output_mode: BuiltinOutputMode::Fixed,
334    completion_policy: BuiltinCompletionPolicy::Public,
335    errors: &INTERP2_ERRORS,
336};
337
338fn interp2_error_with_message(
339    message: impl Into<String>,
340    error: &'static BuiltinErrorDescriptor,
341) -> RuntimeError {
342    let mut builder = build_runtime_error(message).with_builtin(NAME);
343    if let Some(identifier) = error.identifier {
344        builder = builder.with_identifier(identifier);
345    }
346    builder.build()
347}
348
349fn interp2_invalid_argument(detail: impl AsRef<str>) -> RuntimeError {
350    interp2_error_with_message(
351        format!(
352            "{}: {}",
353            INTERP2_ERROR_INVALID_ARGUMENT.message,
354            detail.as_ref()
355        ),
356        &INTERP2_ERROR_INVALID_ARGUMENT,
357    )
358}
359
360fn interp2_invalid_input(detail: impl AsRef<str>) -> RuntimeError {
361    interp2_error_with_message(
362        format!(
363            "{}: {}",
364            INTERP2_ERROR_INVALID_INPUT.message,
365            detail.as_ref()
366        ),
367        &INTERP2_ERROR_INVALID_INPUT,
368    )
369}
370
371fn interp2_map_error(err: RuntimeError, fallback: &'static BuiltinErrorDescriptor) -> RuntimeError {
372    if err.identifier().is_some() {
373        err
374    } else {
375        interp2_error_with_message(err.message().to_string(), fallback)
376    }
377}
378
379#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::math::interpolation::interp2")]
380pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
381    name: NAME,
382    op_kind: GpuOpKind::Custom("interpolation-2d"),
383    supported_precisions: &[ScalarType::F32, ScalarType::F64],
384    broadcast: BroadcastSemantics::Matlab,
385    provider_hooks: &[],
386    constant_strategy: ConstantStrategy::InlineLiteral,
387    residency: ResidencyPolicy::GatherImmediately,
388    nan_mode: ReductionNaN::Include,
389    two_pass_threshold: None,
390    workgroup_size: None,
391    accepts_nan_mode: false,
392    notes: "Initial implementation gathers GPU inputs to the CPU reference path. Bilinear and nearest kernels are good future provider candidates.",
393};
394
395#[runmat_macros::register_fusion_spec(
396    builtin_path = "crate::builtins::math::interpolation::interp2"
397)]
398pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
399    name: NAME,
400    shape: ShapeRequirements::Any,
401    constant_strategy: ConstantStrategy::InlineLiteral,
402    elementwise: None,
403    reduction: None,
404    emits_nan: true,
405    notes: "interp2 is currently a runtime sink.",
406};
407
408fn interp2_type(args: &[Type], _ctx: &ResolveContext) -> Type {
409    let query = match args.len() {
410        0..=2 => return Type::tensor(),
411        3 | 4 => args.get(1),
412        _ => args.get(3),
413    };
414    match query {
415        Some(Type::Num | Type::Int | Type::Bool) => Type::Num,
416        Some(Type::Tensor { shape }) | Some(Type::Logical { shape }) => Type::Tensor {
417            shape: shape.clone(),
418        },
419        _ => Type::tensor(),
420    }
421}
422
423#[runtime_builtin(
424    name = "interp2",
425    category = "math/interpolation",
426    summary = "Interpolate two-dimensional gridded data.",
427    keywords = "interp2,interpolation,bilinear,nearest,grid,meshgrid",
428    accel = "sink",
429    sink = true,
430    type_resolver(interp2_type),
431    descriptor(crate::builtins::math::interpolation::interp2::INTERP2_DESCRIPTOR),
432    builtin_path = "crate::builtins::math::interpolation::interp2"
433)]
434async fn interp2_builtin(args: Vec<Value>) -> crate::BuiltinResult<Value> {
435    let parsed = ParsedInterp2::parse(args)
436        .await
437        .map_err(|err| interp2_map_error(err, &INTERP2_ERROR_INVALID_INPUT))?;
438    let data =
439        evaluate_grid(&parsed).map_err(|err| interp2_map_error(err, &INTERP2_ERROR_INTERNAL))?;
440    if data.len() == 1 {
441        return Ok(Value::Num(data[0]));
442    }
443    let tensor = Tensor::new(data, parsed.output_shape).map_err(|err| {
444        interp2_error_with_message(format!("{NAME}: {err}"), &INTERP2_ERROR_INTERNAL)
445    })?;
446    Ok(Value::Tensor(tensor))
447}
448
449struct ParsedInterp2 {
450    x_axis: Vec<f64>,
451    y_axis: Vec<f64>,
452    z: Tensor,
453    xq: Vec<f64>,
454    yq: Vec<f64>,
455    output_shape: Vec<usize>,
456    method: InterpMethod,
457    extrap: Extrapolation,
458}
459
460impl ParsedInterp2 {
461    async fn parse(args: Vec<Value>) -> crate::BuiltinResult<Self> {
462        if args.len() < 3 {
463            return Err(interp2_invalid_argument(
464                "expected Z, Xq, and Yq or X, Y, Z, Xq, and Yq",
465            ));
466        }
467
468        let mut method = InterpMethod::Linear;
469        let mut extrap = Extrapolation::Nan;
470        let explicit_axes = args.len() >= 5 && !is_option_arg(&args[3]);
471        let (x_axis, y_axis, z, xq_value, yq_value, options) = if explicit_axes {
472            let mut iter = args.into_iter();
473            let x = iter.next().expect("X");
474            let y = iter.next().expect("Y");
475            let z_value = iter.next().expect("Z");
476            let z = z_tensor(z_value).await?;
477            let (x_axis, y_axis) = axes_from_values(x, y, z.rows, z.cols).await?;
478            let xq = iter.next().expect("Xq");
479            let yq = iter.next().expect("Yq");
480            (x_axis, y_axis, z, xq, yq, iter.collect::<Vec<_>>())
481        } else {
482            let mut iter = args.into_iter();
483            let z_value = iter.next().expect("Z");
484            let z = z_tensor(z_value).await?;
485            let x_axis: Vec<f64> = (1..=z.cols).map(|v| v as f64).collect();
486            let y_axis: Vec<f64> = (1..=z.rows).map(|v| v as f64).collect();
487            let xq = iter.next().expect("Xq");
488            let yq = iter.next().expect("Yq");
489            (x_axis, y_axis, z, xq, yq, iter.collect::<Vec<_>>())
490        };
491
492        validate_axis(&x_axis, "X")?;
493        validate_axis(&y_axis, "Y")?;
494        let xq = query_points(xq_value, NAME).await?;
495        let yq = query_points(yq_value, NAME).await?;
496        let (xq_values, yq_values, output_shape) = align_queries(xq, yq)?;
497
498        for option in &options {
499            if let Some(parsed) = parse_extrapolation(option, NAME).await? {
500                extrap = parsed;
501                continue;
502            }
503            if let Some(parsed) = parse_method(option, NAME)? {
504                match parsed {
505                    InterpMethod::Linear | InterpMethod::Nearest => method = parsed,
506                    _ => {
507                        return Err(interp2_invalid_argument(
508                            "only linear and nearest methods are supported",
509                        ));
510                    }
511                }
512                continue;
513            }
514            return Err(interp2_error_with_message(
515                "interp2: unsupported interpolation option",
516                &INTERP2_ERROR_INVALID_ARGUMENT,
517            ));
518        }
519
520        Ok(Self {
521            x_axis,
522            y_axis,
523            z,
524            xq: xq_values,
525            yq: yq_values,
526            output_shape,
527            method,
528            extrap,
529        })
530    }
531}
532
533fn is_option_arg(value: &Value) -> bool {
534    crate::builtins::common::random_args::keyword_of(value).is_some()
535}
536
537async fn z_tensor(value: Value) -> crate::BuiltinResult<Tensor> {
538    let gathered = dispatcher::gather_if_needed_async(&value).await?;
539    let z =
540        tensor::value_into_tensor_for(NAME, gathered).map_err(|err| interp2_invalid_input(&err))?;
541    if z.shape.len() > 2 {
542        return Err(interp2_invalid_argument("Z must be a 2-D matrix"));
543    }
544    if z.rows < 2 || z.cols < 2 {
545        return Err(interp2_invalid_argument(
546            "Z must have at least two rows and two columns",
547        ));
548    }
549    Ok(z)
550}
551
552async fn axes_from_values(
553    x: Value,
554    y: Value,
555    rows: usize,
556    cols: usize,
557) -> crate::BuiltinResult<(Vec<f64>, Vec<f64>)> {
558    let x_axis = axis_from_value(x, rows, cols, true).await?;
559    let y_axis = axis_from_value(y, rows, cols, false).await?;
560    Ok((x_axis, y_axis))
561}
562
563async fn axis_from_value(
564    value: Value,
565    rows: usize,
566    cols: usize,
567    is_x: bool,
568) -> crate::BuiltinResult<Vec<f64>> {
569    let gathered = dispatcher::gather_if_needed_async(&value).await?;
570    let tensor_value = tensor::value_into_tensor_for(NAME, gathered.clone());
571    if let Ok(t) = tensor_value {
572        if is_vector_shape(&t.shape) {
573            let expected = if is_x { cols } else { rows };
574            if t.data.len() != expected {
575                return Err(interp2_invalid_argument(
576                    "axis vector length must match Z dimensions",
577                ));
578            }
579            return Ok(t.data);
580        }
581        if t.rows == rows && t.cols == cols {
582            return if is_x {
583                Ok((0..cols).map(|col| t.data[col * rows]).collect())
584            } else {
585                Ok((0..rows).map(|row| t.data[row]).collect())
586            };
587        }
588    }
589    let label = if is_x { "X" } else { "Y" };
590    vector_from_value(gathered, label, NAME).await
591}
592
593fn validate_axis(axis: &[f64], label: &str) -> crate::BuiltinResult<()> {
594    if axis.len() < 2 {
595        return Err(interp2_invalid_argument(format!(
596            "{label} axis must contain at least two points"
597        )));
598    }
599    if axis.iter().any(|v| !v.is_finite()) {
600        return Err(interp2_invalid_argument(format!(
601            "{label} axis must be finite"
602        )));
603    }
604    for pair in axis.windows(2) {
605        if pair[1] <= pair[0] {
606            return Err(interp2_invalid_argument(format!(
607                "{label} axis must be strictly increasing"
608            )));
609        }
610    }
611    Ok(())
612}
613
614fn align_queries(
615    xq: super::pp::QueryPoints,
616    yq: super::pp::QueryPoints,
617) -> crate::BuiltinResult<(Vec<f64>, Vec<f64>, Vec<usize>)> {
618    match (xq.values.len(), yq.values.len()) {
619        (1, 1) => Ok((xq.values, yq.values, vec![1, 1])),
620        (1, len) => Ok((vec![xq.values[0]; len], yq.values, yq.shape)),
621        (len, 1) => Ok((xq.values, vec![yq.values[0]; len], xq.shape)),
622        (left, right) if left == right && xq.shape == yq.shape => {
623            Ok((xq.values, yq.values, xq.shape))
624        }
625        _ => Err(interp2_invalid_argument(
626            "Xq and Yq must be scalar or matching-size arrays",
627        )),
628    }
629}
630
631fn evaluate_grid(parsed: &ParsedInterp2) -> crate::BuiltinResult<Vec<f64>> {
632    let mut out = Vec::with_capacity(parsed.xq.len());
633    for (&xq, &yq) in parsed.xq.iter().zip(parsed.yq.iter()) {
634        let value = match parsed.method {
635            InterpMethod::Linear => eval_bilinear(parsed, xq, yq),
636            InterpMethod::Nearest => eval_nearest(parsed, xq, yq),
637            _ => unreachable!("interp2 parse rejects cubic methods"),
638        };
639        out.push(value);
640    }
641    Ok(out)
642}
643
644fn eval_bilinear(parsed: &ParsedInterp2, xq: f64, yq: f64) -> f64 {
645    if !xq.is_finite() || !yq.is_finite() {
646        return f64::NAN;
647    }
648    let allow = matches!(parsed.extrap, Extrapolation::Extrapolate);
649    let Some(col) = interval_index(&parsed.x_axis, xq, allow) else {
650        return out_of_range_value(&parsed.extrap);
651    };
652    let Some(row) = interval_index(&parsed.y_axis, yq, allow) else {
653        return out_of_range_value(&parsed.extrap);
654    };
655    let x0 = parsed.x_axis[col];
656    let x1 = parsed.x_axis[col + 1];
657    let y0 = parsed.y_axis[row];
658    let y1 = parsed.y_axis[row + 1];
659    let tx = (xq - x0) / (x1 - x0);
660    let ty = (yq - y0) / (y1 - y0);
661    let z00 = z_at(&parsed.z, row, col);
662    let z10 = z_at(&parsed.z, row, col + 1);
663    let z01 = z_at(&parsed.z, row + 1, col);
664    let z11 = z_at(&parsed.z, row + 1, col + 1);
665    (1.0 - tx) * (1.0 - ty) * z00 + tx * (1.0 - ty) * z10 + (1.0 - tx) * ty * z01 + tx * ty * z11
666}
667
668fn eval_nearest(parsed: &ParsedInterp2, xq: f64, yq: f64) -> f64 {
669    if !xq.is_finite() || !yq.is_finite() {
670        return f64::NAN;
671    }
672    let Some(col) = nearest_index(&parsed.x_axis, xq, &parsed.extrap) else {
673        return out_of_range_value(&parsed.extrap);
674    };
675    let Some(row) = nearest_index(&parsed.y_axis, yq, &parsed.extrap) else {
676        return out_of_range_value(&parsed.extrap);
677    };
678    z_at(&parsed.z, row, col)
679}
680
681fn z_at(z: &Tensor, row: usize, col: usize) -> f64 {
682    z.data[row + col * z.rows]
683}
684
685fn nearest_index(axis: &[f64], q: f64, extrap: &Extrapolation) -> Option<usize> {
686    if q < axis[0] {
687        return matches!(extrap, Extrapolation::Extrapolate).then_some(0);
688    }
689    let last = axis.len() - 1;
690    if q > axis[last] {
691        return matches!(extrap, Extrapolation::Extrapolate).then_some(last);
692    }
693    match axis.binary_search_by(|probe| probe.partial_cmp(&q).unwrap()) {
694        Ok(index) => Some(index),
695        Err(index) => {
696            let left = index.saturating_sub(1);
697            let right = index.min(last);
698            if (q - axis[left]).abs() <= (axis[right] - q).abs() {
699                Some(left)
700            } else {
701                Some(right)
702            }
703        }
704    }
705}
706
707#[cfg(test)]
708mod tests {
709    use super::*;
710    use futures::executor::block_on;
711
712    fn row(values: &[f64]) -> Value {
713        Value::Tensor(Tensor::new(values.to_vec(), vec![1, values.len()]).expect("tensor"))
714    }
715
716    #[test]
717    fn interp2_implicit_axes_bilinear_scalar() {
718        let z = Value::Tensor(Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).expect("tensor"));
719        let value =
720            block_on(interp2_builtin(vec![z, Value::Num(1.5), Value::Num(1.5)])).expect("interp2");
721        let Value::Num(result) = value else {
722            panic!("expected scalar");
723        };
724        assert!((result - 2.5).abs() < 1e-12);
725    }
726
727    #[test]
728    fn interp2_vector_axes_nearest() {
729        let z = Value::Tensor(Tensor::new(vec![1.0, 3.0, 2.0, 4.0], vec![2, 2]).expect("tensor"));
730        let value = block_on(interp2_builtin(vec![
731            row(&[10.0, 20.0]),
732            row(&[100.0, 200.0]),
733            z,
734            Value::Num(18.0),
735            Value::Num(120.0),
736            Value::String("nearest".to_string()),
737        ]))
738        .expect("interp2");
739        assert_eq!(value, Value::Num(2.0));
740    }
741
742    #[test]
743    fn interp2_descriptor_signatures_cover_surface() {
744        let labels: Vec<&str> = INTERP2_DESCRIPTOR
745            .signatures
746            .iter()
747            .map(|signature| signature.label)
748            .collect();
749        assert!(labels.contains(&"Vq = interp2(Z, Xq, Yq)"));
750        assert!(labels.contains(&"Vq = interp2(X, Y, Z, Xq, Yq)"));
751        assert!(labels.contains(&"Vq = interp2(X, Y, Z, Xq, Yq, method, extrap)"));
752    }
753
754    #[test]
755    fn interp2_descriptor_errors_have_stable_codes() {
756        let codes: Vec<&str> = INTERP2_DESCRIPTOR
757            .errors
758            .iter()
759            .map(|error| error.code)
760            .collect();
761        assert!(codes.contains(&"RM.INTERP2.INVALID_ARGUMENT"));
762        assert!(codes.contains(&"RM.INTERP2.INVALID_INPUT"));
763        assert!(codes.contains(&"RM.INTERP2.INTERNAL"));
764    }
765
766    #[test]
767    fn interp2_too_few_args_uses_stable_identifier() {
768        let err = block_on(interp2_builtin(vec![Value::Num(1.0), Value::Num(2.0)]))
769            .expect_err("expected interp2 argument error");
770        assert_eq!(err.identifier(), INTERP2_ERROR_INVALID_ARGUMENT.identifier);
771    }
772}