135 lines
4.2 KiB
Go
135 lines
4.2 KiB
Go
// 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
|
|||
|
|
}
|