Files
tensor/linalg/spline.go
T

135 lines
4.2 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 linalg
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import "math"
// Natural cubic spline interpolation. Given n points (xᵢ, yᵢ) with
// strictly increasing x, the natural spline is the unique piecewise
// cubic that passes through every point, is C² across the interior
// knots, and has zero second derivative at both ends. The second
// derivatives come from an n−2 tridiagonal system solved by the
// library's `SolveTridiagonal`, so the setup reuses the sparse
// machinery rather than a second LU.
// CubicSpline holds the natural cubic spline interpolant over sorted
// knots. At returns the interpolated value; Evaluate maps an array of
// query points element-wise.
type CubicSpline struct {
xs, ys, m []float64 // knots, values, second derivatives
n int
}
// NewCubicSpline builds a natural cubic spline through the given
// points. The abscissae must be strictly increasing; at least 3
// points are needed for a piecewise cubic to have interior knots.
func NewCubicSpline(xs, ys *core.Array) (*CubicSpline, error) {
n := xs.Len()
if err := requireReal("NewCubicSpline", xs, ys); err != nil {
return nil, err
}
if xs.NDim() != 1 || ys.NDim() != 1 {
return nil, base.Errf("NewCubicSpline: xs and ys must be 1-D, got shapes %s and %s",
base.ShapeText(xs.Shape()), base.ShapeText(ys.Shape()))
}
if n < 3 {
return nil, base.Errf("NewCubicSpline: at least 3 points are needed, got %d", n)
}
if ys.Len() != n {
return nil, base.Errf("NewCubicSpline: xs has %d points but ys has %d", n, ys.Len())
}
knots := make([]float64, n)
for i := range n {
knots[i] = xs.FloatAt(i)
if i > 0 && knots[i] <= knots[i-1] {
return nil, base.Errf("NewCubicSpline: abscissae must be strictly increasing, found %v after %v",
knots[i], knots[i-1])
}
}
vals := make([]float64, n)
for i := range n {
vals[i] = ys.FloatAt(i)
}
// Second derivatives via the tridiagonal system for the interior
// knots, with natural (zero) boundary conditions.
h := make([]float64, n-1)
for i := range n - 1 {
h[i] = knots[i+1] - knots[i]
}
m := n - 2 // interior knots: the tridiagonal system is m×m
sub := make([]float64, m-1)
main := make([]float64, m)
sup := make([]float64, m-1)
rhs := make([]float64, m)
for i := range m {
main[i] = 2 * (h[i] + h[i+1])
if i > 0 {
sub[i-1] = h[i]
}
if i < m-1 {
sup[i] = h[i+1]
}
slope1 := (vals[i+1] - vals[i]) / h[i]
slope2 := (vals[i+2] - vals[i+1]) / h[i+1]
rhs[i] = 6 * (slope2 - slope1)
}
subArr := floatsToArray(sub, []int{m - 1})
mainArr := floatsToArray(main, []int{m})
supArr := floatsToArray(sup, []int{m - 1})
rhsArr := floatsToArray(rhs, []int{m})
mArr, err := SolveTridiagonal(subArr, mainArr, supArr, rhsArr)
if err != nil {
return nil, base.Errf("NewCubicSpline: %w", err)
}
sec := make([]float64, n)
sec[0], sec[n-1] = 0, 0
for i := 1; i < n-1; i++ {
sec[i] = mArr.FloatAt(i - 1)
}
return &CubicSpline{xs: knots, ys: vals, m: sec, n: n}, nil
}
// At evaluates the spline at one query point. Values outside the
// knot range return NaN, mirroring the convention that extrapolation
// without a boundary condition is undefined.
func (s *CubicSpline) At(x float64) float64 {
if x < s.xs[0] || x > s.xs[s.n-1] {
return math.NaN()
}
// Binary search for the interval [x_lo, x_lo+1].
lo, hi := 0, s.n-1
for hi-lo > 1 {
mid := (lo + hi) / 2
if s.xs[mid] <= x {
lo = mid
} else {
hi = mid
}
}
h := s.xs[lo+1] - s.xs[lo]
a := (s.xs[lo+1] - x) / h
b := (x - s.xs[lo]) / h
return a*s.ys[lo] + b*s.ys[lo+1] +
((a*a*a-a)*s.m[lo]+(b*b*b-b)*s.m[lo+1])*h*h/6
}
// Evaluate maps an array of query points through the spline. Like
// NewCubicSpline it reads its argument as real numbers, so a complex
// query array is refused.
func (s *CubicSpline) Evaluate(x *core.Array) (*core.Array, error) {
if err := requireReal("CubicSpline.Evaluate", x); err != nil {
return nil, err
}
out := core.New(core.Float, append([]int{}, x.Shape()...)...)
for i := range x.Len() {
out.RawFloats()[i] = s.At(x.FloatAt(i))
}
return out, nil
}