200 lines
7.9 KiB
Go
200 lines
7.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package core
|
||
|
||
import (
|
||
"math"
|
||
"sync"
|
||
)
|
||
|
||
// Bilinear interpolation over a regular grid, the two-dimensional
|
||
// companion of the 1-D Interpolate in arrayutil.go. Where the 1-D
|
||
// entry point clamps queries to the boundary, the grid version
|
||
// refuses them: in two dimensions a silently clamped value is far
|
||
// harder to spot, and the honest answer to an out-of-domain query is
|
||
// an error.
|
||
|
||
// interp2dMinPerWorker is the per-worker chunk floor for the bilinear
|
||
// queries. A query costs a couple of Floor calls and a dozen flops
|
||
// over four grid reads, roughly a dozen nanoseconds, so a chunk in
|
||
// the hundreds already outweighs a worker's start-up; the sweep around
|
||
// the value pinned 256 ahead of both 1024 and 4096 at the
|
||
// million-query size.
|
||
const interp2dMinPerWorker = 256
|
||
|
||
// Interpolate2D evaluates the bilinear interpolation of a regular
|
||
// grid at a set of query points. grid is a rank-2, rows × cols array
|
||
// whose element (i, j) samples the interpolated field at
|
||
// (x0 + j·dx, y0 + i·dy); xs and ys hold the query coordinates, must
|
||
// share their length, and the result takes xs's shape. Bilinear
|
||
// interpolation is exact on every function that is bilinear within a
|
||
// cell, so linear fields come back unchanged and grid nodes are
|
||
// reproduced exactly.
|
||
//
|
||
// Every query must lie inside the closed rectangle the grid spans,
|
||
// boundaries included; a query outside it is an error, never a
|
||
// silent clamp. Either sign of dx and dy works, so descending axes
|
||
// need no preprocessing.
|
||
//
|
||
// Errors: grid not rank 2 or an extent below 2, complex inputs, a
|
||
// grid dtype the surface refuses (bool and the narrow integer
|
||
// widths, refused by name; convert with Astype first), xs and ys of
|
||
// different lengths, a zero or non-finite spacing or origin
|
||
// coordinate, a non-finite query coordinate, or a query outside the
|
||
// domain.
|
||
func Interpolate2D(grid, xs, ys *Array, x0, y0, dx, dy float64) (*Array, error) {
|
||
if grid.NDim() != 2 {
|
||
return nil, errf("Interpolate2D: grid must be rank 2, got shape %s", shapeText(grid.Shape()))
|
||
}
|
||
rows, cols := grid.Shape()[0], grid.Shape()[1]
|
||
if rows < 2 || cols < 2 {
|
||
return nil, errf("Interpolate2D: grid needs at least 2×2 samples, got %d×%d", rows, cols)
|
||
}
|
||
if grid.dt == Complex {
|
||
return nil, errf("Interpolate2D: complex grids are not supported")
|
||
}
|
||
if xs.Len() != ys.Len() {
|
||
return nil, errf("Interpolate2D: xs and ys must share their length, got %d and %d", xs.Len(), ys.Len())
|
||
}
|
||
if xs.dt == Complex || ys.dt == Complex {
|
||
return nil, errf("Interpolate2D: complex query coordinates are not supported")
|
||
}
|
||
if narrowRefused(grid.dt) {
|
||
// Bool and the narrow integer grids carry no bilinear kernel,
|
||
// and the default grid reads below would reach a nil float64
|
||
// payload for them: the narrow-dtype refusal names the grid dtype
|
||
// and the conversion instead, after the rank and complex gates
|
||
// whose texts the existing pins carry.
|
||
return nil, errf("Interpolate2D: dtype %s is not supported; convert with Astype", grid.dt)
|
||
}
|
||
if dx == 0 || dy == 0 || math.IsNaN(dx) || math.IsNaN(dy) ||
|
||
math.IsInf(dx, 0) || math.IsInf(dy, 0) ||
|
||
math.IsNaN(x0) || math.IsNaN(y0) || math.IsInf(x0, 0) || math.IsInf(y0, 0) {
|
||
return nil, errf("Interpolate2D: the origin and spacings must be finite with dx, dy non-zero, got x0=%g y0=%g dx=%g dy=%g", x0, y0, dx, dy)
|
||
}
|
||
// The domain runs between the first and last sample along each
|
||
// axis, which for a negative spacing means a reversed interval.
|
||
xLo, xHi := min(x0, x0+float64(cols-1)*dx), max(x0, x0+float64(cols-1)*dx)
|
||
yLo, yHi := min(y0, y0+float64(rows-1)*dy), max(y0, y0+float64(rows-1)*dy)
|
||
// The coordinates and the grid are read through dense payload
|
||
// windows: materialise once, widen the coordinates exactly as
|
||
// floatAt does (float32 and int convert exactly), and bind the
|
||
// grid slices so the loop pays no per-element accessor dispatch.
|
||
xf, yf := widened(xs), widened(ys)
|
||
g := grid.materialise()
|
||
n := xs.Len()
|
||
out := &Array{shape: append([]int{}, xs.shape...), dt: Float}
|
||
out.alloc(n)
|
||
// The workers write disjoint output slots and report a query
|
||
// error through a mutex touched only on the failing path; the
|
||
// sequential contract is that the first offending query wins, so
|
||
// the reported error is the one with the smallest query index.
|
||
var mu sync.Mutex
|
||
var firstErr error
|
||
firstIdx := 0
|
||
fail := func(i int, err error) {
|
||
mu.Lock()
|
||
defer mu.Unlock()
|
||
if firstErr == nil || i < firstIdx {
|
||
firstErr, firstIdx = err, i
|
||
}
|
||
}
|
||
parallelMin(n, interp2dMinPerWorker, func(s, e int) {
|
||
for i := s; i < e; i++ {
|
||
x, y := xf[i], yf[i]
|
||
if math.IsNaN(x) || math.IsInf(x, 0) || math.IsNaN(y) || math.IsInf(y, 0) {
|
||
fail(i, errf("Interpolate2D: query %d is not finite (x=%g, y=%g)", i, x, y))
|
||
return
|
||
}
|
||
if x < xLo || x > xHi || y < yLo || y > yHi {
|
||
fail(i, errf("Interpolate2D: query (%g, %g) lies outside the grid domain x ∈ [%g, %g], y ∈ [%g, %g]", x, y, xLo, xHi, yLo, yHi))
|
||
return
|
||
}
|
||
// Cell indices from the fractional grid position, folded onto
|
||
// the final cell at the exact upper boundary; the lower clamps
|
||
// only absorb rounding at the exact lower boundary.
|
||
gx := (x - x0) / dx
|
||
gy := (y - y0) / dy
|
||
j := min(max(int(math.Floor(gx)), 0), cols-2)
|
||
r := min(max(int(math.Floor(gy)), 0), rows-2)
|
||
tx := gx - float64(j)
|
||
ty := gy - float64(r)
|
||
var v00, v01, v10, v11 float64
|
||
switch g.dt {
|
||
case Float16:
|
||
// The widening of a half grid value is exact, so the
|
||
// fold runs on the same bits the float32 path runs on.
|
||
b := r*cols + j
|
||
v00, v01 = HalfToFloat64(g.halves[b]), HalfToFloat64(g.halves[b+1])
|
||
v10, v11 = HalfToFloat64(g.halves[b+cols]), HalfToFloat64(g.halves[b+cols+1])
|
||
case Float32:
|
||
b := r*cols + j
|
||
v00, v01 = float64(g.floats32[b]), float64(g.floats32[b+1])
|
||
v10, v11 = float64(g.floats32[b+cols]), float64(g.floats32[b+cols+1])
|
||
case Int:
|
||
b := r*cols + j
|
||
v00, v01 = float64(g.ints[b]), float64(g.ints[b+1])
|
||
v10, v11 = float64(g.ints[b+cols]), float64(g.ints[b+cols+1])
|
||
default:
|
||
b := r*cols + j
|
||
v00, v01 = g.floats[b], g.floats[b+1]
|
||
v10, v11 = g.floats[b+cols], g.floats[b+cols+1]
|
||
}
|
||
// The explicit conversions are the FMA fence, the one the
|
||
// window catalogue uses: a contracting build would fuse the
|
||
// bare products into the sums and the answer would drift a
|
||
// ulp off the portable bits.
|
||
top := float64((1-tx)*v00) + float64(tx*v01)
|
||
bot := float64((1-tx)*v10) + float64(tx*v11)
|
||
out.floats[i] = float64((1-ty)*top) + float64(ty*bot)
|
||
}
|
||
})
|
||
if firstErr != nil {
|
||
return nil, firstErr
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// widened returns the array's elements as a float64 payload, widened
|
||
// exactly as floatAt widens: float16, float32, int, bool and the narrow
|
||
// integer widths all convert exactly, and float aliases the payload
|
||
// itself. The loops are bounded by the logical length: a rebased view
|
||
// carries a payload longer than its extent.
|
||
func widened(a *Array) []float64 {
|
||
m := a.materialise()
|
||
switch m.dt {
|
||
case Float16:
|
||
out := make([]float64, m.Len())
|
||
for i := range m.Len() {
|
||
out[i] = HalfToFloat64(m.halves[i])
|
||
}
|
||
return out
|
||
case Float32:
|
||
out := make([]float64, m.Len())
|
||
for i := range m.Len() {
|
||
out[i] = float64(m.floats32[i])
|
||
}
|
||
return out
|
||
case Int:
|
||
out := make([]float64, m.Len())
|
||
for i := range m.Len() {
|
||
out[i] = float64(m.ints[i])
|
||
}
|
||
return out
|
||
case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32:
|
||
// The narrow real class widens through the same accessor the
|
||
// query gates let through: bool reads 0/1 and every narrow
|
||
// integer widens exactly, the values floatAt hands over. Without
|
||
// this arm a bool query array reached the default and read a nil
|
||
// float64 payload, panicking on the first slot.
|
||
out := make([]float64, m.Len())
|
||
for i := range m.Len() {
|
||
out[i] = m.floatAt(i)
|
||
}
|
||
return out
|
||
default:
|
||
return m.floats
|
||
}
|
||
}
|