Files
tensor/linalg/tridiag.go
T

142 lines
5.1 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"
)
// 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
}