Files
tensor/internal/core/interpolate2d.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

200 lines
7.9 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"
"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
}
}