Files
tensor/internal/core/interpolate2d_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

234 lines
8.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}
}
}