142 lines
5.1 KiB
Go
142 lines
5.1 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"
|
||
)
|
||
|
||
// Tridiagonal solvers. The Thomas algorithm and its periodic
|
||
// (Sherman-Morrison) variant solve the tridiagonal systems that
|
||
// discretised one-dimensional differential operators produce, in O(n)
|
||
// time and O(n) memory, versus the O(n³) dense `Solve`. All inputs
|
||
// are rank-1 vectors; the sub-, main- and superdiagonals are separate
|
||
// vectors of lengths n−1, n and n−1.
|
||
|
||
// SolveTridiagonal returns the vector x solving the tridiagonal system
|
||
// with lower diagonal a (length n−1), main diagonal b (length n),
|
||
// upper diagonal c (length n−1) and right-hand side d (length n),
|
||
// by the Thomas algorithm. A zero pivot is refused.
|
||
func SolveTridiagonal(a, b, c, d *core.Array) (*core.Array, error) {
|
||
n := b.Len()
|
||
if err := requireReal("SolveTridiagonal", a, b, c, d); err != nil {
|
||
return nil, err
|
||
}
|
||
if n == 0 {
|
||
return nil, base.Errf("SolveTridiagonal: empty system")
|
||
}
|
||
if a.Len() != n-1 || c.Len() != n-1 || d.Len() != n {
|
||
return nil, base.Errf("SolveTridiagonal: diagonal lengths %d/%d/%d do not match rhs %d",
|
||
a.Len(), b.Len(), c.Len(), d.Len())
|
||
}
|
||
// The elimination is internal/base's shared kernel, the same one
|
||
// the integrate package's PDE sweeps run against their scratch. A
|
||
// dense float64 operand hands the kernel its payload directly, the
|
||
// same values FloatAt read; every other real dtype reaches the
|
||
// kernel through the FloatAt promotion the element-wise walk
|
||
// applied.
|
||
av, bv, cv, dv := triDiagonals(a, b, c, d, n)
|
||
cp := make([]float64, n)
|
||
dp := make([]float64, n)
|
||
x := make([]float64, n)
|
||
if err := base.TriSolve(x, cp, dp, av, bv, cv, dv); err != nil {
|
||
return nil, err
|
||
}
|
||
return floatsToArray(x, []int{n}), nil
|
||
}
|
||
|
||
// triDiagonals materialises the four dense float64 operands the shared
|
||
// kernel reads. A dense float64 array's payload slice is exactly its
|
||
// visible range, so it is passed through untouched; any other real
|
||
// dtype is promoted element by element, the conversion FloatAt
|
||
// performs. A strided operand takes the promotion path, where the
|
||
// accessor walks the logical order.
|
||
func triDiagonals(a, b, c, d *core.Array, n int) (av, bv, cv, dv []float64) {
|
||
extract := func(arr *core.Array, m int) []float64 {
|
||
if arr.Dtype() == core.Float && !arr.Strided() {
|
||
return arr.RawFloats()[:m]
|
||
}
|
||
out := make([]float64, m)
|
||
for i := range m {
|
||
out[i] = arr.FloatAt(i)
|
||
}
|
||
return out
|
||
}
|
||
return extract(a, n-1), extract(b, n), extract(c, n-1), extract(d, n)
|
||
}
|
||
|
||
// SolveCyclicTridiagonal returns the vector x solving the periodic
|
||
// tridiagonal system whose corner entries couple the ends: a[0] is
|
||
// the wrap corner A[0][n−1], c[n−1] is the wrap corner A[n−1][0], and
|
||
// a[1..n−1], b[0..n−1], c[0..n−2] fill the interior bands. Built on
|
||
// the Sherman-Morrison rank-one update of the Thomas elimination,
|
||
// still O(n).
|
||
func SolveCyclicTridiagonal(a, b, c, d *core.Array) (*core.Array, error) {
|
||
n := b.Len()
|
||
if err := requireReal("SolveCyclicTridiagonal", a, b, c, d); err != nil {
|
||
return nil, err
|
||
}
|
||
if n < 3 {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: needs at least 3 unknowns, got %d", n)
|
||
}
|
||
if a.Len() != n || c.Len() != n || d.Len() != n {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: diagonal lengths %d/%d/%d do not match rhs %d",
|
||
a.Len(), b.Len(), c.Len(), d.Len())
|
||
}
|
||
alpha := a.FloatAt(0)
|
||
beta := c.FloatAt(n - 1)
|
||
gamma := -b.FloatAt(0)
|
||
if gamma == 0 {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: gamma = −b[0] vanishes, system is singular")
|
||
}
|
||
|
||
// Strip the rank-one correction: the interior tridiagonal T has
|
||
// b′[0] = b[0] − γ, b′[n−1] = b[n−1] − αβ/γ, no wrap corners.
|
||
bb := make([]float64, n)
|
||
for i := range n {
|
||
bb[i] = b.FloatAt(i)
|
||
}
|
||
bb[0] -= gamma
|
||
bb[n-1] -= alpha * beta / gamma
|
||
aInt := make([]float64, n-1)
|
||
cInt := make([]float64, n-1)
|
||
for i := 1; i < n; i++ {
|
||
aInt[i-1] = a.FloatAt(i)
|
||
}
|
||
for i := range n - 1 {
|
||
cInt[i] = c.FloatAt(i)
|
||
}
|
||
aArr := floatsToArray(aInt, []int{n - 1})
|
||
cArr := floatsToArray(cInt, []int{n - 1})
|
||
bArr := floatsToArray(bb, []int{n})
|
||
|
||
// Two Thomas solves against the modified matrix: one for the
|
||
// right-hand side, one for the correction vector u.
|
||
ym, err := SolveTridiagonal(aArr, bArr, cArr, d)
|
||
if err != nil {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: %w", err)
|
||
}
|
||
uVec := make([]float64, n)
|
||
uVec[0] = gamma
|
||
uVec[n-1] = beta
|
||
uArr := floatsToArray(uVec, []int{n})
|
||
zm, err := SolveTridiagonal(aArr, bArr, cArr, uArr)
|
||
if err != nil {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: %w", err)
|
||
}
|
||
// v = (1, 0, …, 0, α/γ): x = y − z·(vᵀy)/(1 + vᵀz).
|
||
vty := ym.FloatAt(0) + alpha/gamma*ym.FloatAt(n-1)
|
||
vtz := zm.FloatAt(0) + alpha/gamma*zm.FloatAt(n-1)
|
||
if 1+vtz == 0 {
|
||
return nil, base.Errf("SolveCyclicTridiagonal: singular Sherman-Morrison denominator")
|
||
}
|
||
fact := vty / (1 + vtz)
|
||
x := make([]float64, n)
|
||
for i := range n {
|
||
x[i] = ym.FloatAt(i) - fact*zm.FloatAt(i)
|
||
}
|
||
return floatsToArray(x, []int{n}), nil
|
||
}
|