Files
tensor/internal/core/interpolate2d_test.go
T

234 lines
8.5 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"strings"
"testing"
)
// TestInterpolate2DLinearField checks that a linear field comes back
// unchanged: bilinear interpolation is exact on it, so the queries
// pin hand-computed values of 2 + 3x − 1.5y on a 5×4 grid that does
// not start at the origin and has unequal spacings.
func TestInterpolate2DLinearField(t *testing.T) {
const x0, dx, y0, dy = 0.0, 0.5, 10.0, 2.0
gridVals := make([]float64, 0, 20)
for i := range 5 {
for j := range 4 {
x := x0 + float64(j)*dx
y := y0 + float64(i)*dy
gridVals = append(gridVals, 2+3*x-1.5*y)
}
}
grid := mustFloats(t, gridVals, 5, 4)
qx := []float64{1.3, 0, 1.5, 0.75, 0.25}
qy := []float64{12.7, 10, 18, 16, 11}
want := []float64{-13.15, -13, -20.5, -19.75, 2 + 0.75 - 16.5}
got, err := Interpolate2D(grid, mustFloats(t, qx), mustFloats(t, qy), x0, y0, dx, dy)
if err != nil {
t.Fatalf("Interpolate2D: %v", err)
}
for i := range want {
if v := got.FloatAt(i); math.Abs(v-want[i]) > 1e-12*math.Abs(want[i]) {
t.Errorf("query (%g, %g) = %v, want %v", qx[i], qy[i], v, want[i])
}
}
}
// TestInterpolate2DPatchMidpoints pins exact values inside a single
// bilinear patch: the cell corners (1, 2; 3, 7) average to 13/4 at
// the centre, and the off-centre point lands on 53/16. Both targets
// are dyadic, so the comparison is exact.
func TestInterpolate2DPatchMidpoints(t *testing.T) {
grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2)
got, err := Interpolate2D(grid, mustFloats(t, []float64{0.5, 0.25}), mustFloats(t, []float64{0.5, 0.75}), 0, 0, 1, 1)
if err != nil {
t.Fatalf("Interpolate2D: %v", err)
}
if got.FloatAt(0) != 3.25 {
t.Errorf("centre = %v, want 3.25", got.FloatAt(0))
}
if got.FloatAt(1) != 3.3125 {
t.Errorf("off-centre = %v, want 3.3125", got.FloatAt(1))
}
}
// TestInterpolate2DNodes checks that every grid node is reproduced
// exactly, and that a grid with a descending x axis (dx < 0) gives
// the mirrored result.
func TestInterpolate2DNodes(t *testing.T) {
grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2)
for i := range 2 {
for j := range 2 {
got, err := Interpolate2D(grid, mustFloats(t, []float64{float64(j)}), mustFloats(t, []float64{float64(i)}), 0, 0, 1, 1)
if err != nil {
t.Fatalf("Interpolate2D(%d, %d): %v", i, j, err)
}
if want := grid.FloatAt(2*i + j); got.FloatAt(0) != want {
t.Errorf("node (%d, %d) = %v, want %v", i, j, got.FloatAt(0), want)
}
}
}
got, err := Interpolate2D(grid, mustFloats(t, []float64{0.5}), mustFloats(t, []float64{0.5}), 1, 0, -1, 1)
if err != nil {
t.Fatalf("Interpolate2D descending: %v", err)
}
if got.FloatAt(0) != 3.25 {
t.Errorf("descending axis: %v, want 3.25", got.FloatAt(0))
}
}
// TestInterpolate2DOutsideRejects pins the no-silent-clamping
// contract: queries beyond any edge of the domain, and non-finite
// coordinates, are errors.
func TestInterpolate2DOutsideRejects(t *testing.T) {
grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2)
cases := []struct {
name string
xs, ys []float64
errSubstr string
}{
{"below x", []float64{-0.1}, []float64{0.5}, "outside the grid domain"},
{"above x", []float64{1.0001}, []float64{0.5}, "outside the grid domain"},
{"below y", []float64{0.5}, []float64{-0.5}, "outside the grid domain"},
{"above y", []float64{0.5}, []float64{1.5}, "outside the grid domain"},
{"not a number", []float64{0.5}, []float64{math.NaN()}, "not finite"},
}
for _, c := range cases {
_, err := Interpolate2D(grid, mustFloats(t, c.xs), mustFloats(t, c.ys), 0, 0, 1, 1)
if err == nil {
t.Errorf("%s: expected an error", c.name)
} else if !strings.Contains(err.Error(), c.errSubstr) {
t.Errorf("%s: error %q lacks %q", c.name, err, c.errSubstr)
}
}
}
// TestInterpolate2DValidation pins the input contracts: rank-2 grids
// of extent ≥ 2, matched coordinate lengths, finite non-zero
// spacings, and real-valued arrays.
func TestInterpolate2DValidation(t *testing.T) {
grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2)
xs := mustFloats(t, []float64{0.5})
ys := mustFloats(t, []float64{0.5})
if _, err := Interpolate2D(mustFloats(t, []float64{1, 2, 3, 4}, 4), xs, ys, 0, 0, 1, 1); err == nil {
t.Error("expected an error for a rank-1 grid")
}
if _, err := Interpolate2D(mustFloats(t, []float64{1, 2, 3, 4}, 1, 4), xs, ys, 0, 0, 1, 1); err == nil {
t.Error("expected an error for a 1×4 grid")
}
if _, err := Interpolate2D(grid, mustFloats(t, []float64{0.5, 0.6}), ys, 0, 0, 1, 1); err == nil {
t.Error("expected an error for mismatched coordinate lengths")
}
if _, err := Interpolate2D(grid, xs, ys, 0, 0, 0, 1); err == nil {
t.Error("expected an error for dx = 0")
}
if _, err := Interpolate2D(grid, xs, ys, math.Inf(1), 0, 1, 1); err == nil {
t.Error("expected an error for a non-finite origin")
}
complexGrid, _ := FromComplexes([]complex128{1, 2, 3, 4}, 2, 2)
if _, err := Interpolate2D(complexGrid, xs, ys, 0, 0, 1, 1); err == nil {
t.Error("expected an error for a complex grid")
}
if _, err := Interpolate2D(grid, mustFromComplexes(t, []complex128{0.5, 1}, 2), ys, 0, 0, 1, 1); err == nil {
t.Error("expected an error for complex query coordinates")
}
}
// TestInterpolate2DShape checks that the result takes the query
// array's shape, not just its length.
func TestInterpolate2DShape(t *testing.T) {
grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2)
xs := mustFloats(t, []float64{0.5, 0.5, 0.5, 0.5}, 2, 2)
ys := mustFloats(t, []float64{0.5, 0.5, 0.5, 0.5})
got, err := Interpolate2D(grid, xs, ys, 0, 0, 1, 1)
if err != nil {
t.Fatalf("Interpolate2D: %v", err)
}
if shape := got.Shape(); len(shape) != 2 || shape[0] != 2 || shape[1] != 2 {
t.Errorf("result shape %v, want (2, 2)", shape)
}
}
// TestInterpolate2DNarrowGridRefused pins the narrow-dtype refusal at
// the entry: bool and the narrow integer grids carry no bilinear
// kernel, and the default grid reads would reach a nil float64 payload
// for them. The refusal names the grid dtype and the conversion, after
// the rank and complex gates whose texts the pins above carry.
func TestInterpolate2DNarrowGridRefused(t *testing.T) {
must := func(a *Array, err error) *Array {
if err != nil {
t.Fatal(err)
}
return a
}
xs := mustFloats(t, []float64{0.5})
ys := mustFloats(t, []float64{0.5})
cases := []struct {
name string
grid *Array
}{
{"bool", must(FromBools([]bool{true, false, true, true}, 2, 2))},
{"int8", must(FromInt8s([]int8{1, 2, 3, 4}, 2, 2))},
{"uint8", must(FromUint8s([]uint8{1, 2, 3, 4}, 2, 2))},
{"int16", must(FromInt16s([]int16{1, 2, 3, 4}, 2, 2))},
{"uint16", must(FromUint16s([]uint16{1, 2, 3, 4}, 2, 2))},
{"int32", must(FromInt32s([]int32{1, 2, 3, 4}, 2, 2))},
{"uint32", must(FromUint32s([]uint32{1, 2, 3, 4}, 2, 2))},
}
for _, c := range cases {
_, err := Interpolate2D(c.grid, xs, ys, 0, 0, 1, 1)
if err == nil || !strings.Contains(err.Error(), "Interpolate2D") ||
!strings.Contains(err.Error(), c.name) ||
!strings.Contains(err.Error(), "convert with Astype") {
t.Errorf("Interpolate2D %s grid: err = %v, want the narrow-dtype refusal naming %s and the conversion",
c.name, err, c.name)
}
}
}
// TestInterpolate2DNarrowQueries pins the narrow query dtypes: a bool or
// narrow integer coordinate array widens exactly the way floatAt widens
// it, answering the values the equivalent float queries answer. The
// widened helper used to read the nil float64 payload for these dtypes,
// so the first query slot panicked instead of interpolating.
func TestInterpolate2DNarrowQueries(t *testing.T) {
g, err := FromFloats([]float64{0, 1, 2, 3}, 2, 2)
if err != nil {
t.Fatal(err)
}
// The field samples 2y + x at (x0+j·dx, y0+i·dy), so the query
// (x, y) answers 2y + x whatever dtype carries the coordinates.
mk := func(a *Array, err error) *Array {
if err != nil {
t.Fatal(err)
}
return a
}
cases := []struct {
name string
xs *Array
}{
{"bool", mk(FromBools([]bool{false, true}, 2))},
{"int8", mk(FromInt8s([]int8{0, 1}, 2))},
{"uint16", mk(FromUint16s([]uint16{0, 1}, 2))},
{"int32", mk(FromInt32s([]int32{0, 1}, 2))},
}
ys := mustFloats(t, []float64{0, 0.5})
for _, c := range cases {
got, err := Interpolate2D(g, c.xs, ys, 0, 0, 1, 1)
if err != nil {
t.Fatalf("Interpolate2D %s queries: %v", c.name, err)
}
want := []float64{0, 2}
for i, w := range want {
if v := got.FloatAt(i); v != w {
t.Errorf("Interpolate2D %s queries[%d] = %v, want %v", c.name, i, v, w)
}
}
}
}