Files
tensor/linalg/tridiag.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

142 lines
5.1 KiB
Go
Raw Permalink 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 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
}