1use 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}