Files
tensor/internal/base/tridiag.go
T

52 lines
1.8 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 base
// The tridiagonal Thomas elimination, one kernel for every caller. The
// linalg package's public SolveTridiagonal and the integrate package's
// PDE step sweeps solve the same systems; both call TriSolve, so the
// arithmetic and the refusal texts exist once. The messages name
// SolveTridiagonal because both public surfaces publish that text
// today, and a message a caller matches is part of the contract.
// TriSolve solves the tridiagonal system with lower diagonal a
// (length n-1), main diagonal b (length n), upper diagonal c (length
// n-1) and right side d (length n), writing the solution into dst. The
// scratch cp and dp must hold at least n elements. A zero pivot is
// refused. Every buffer is fully overwritten before the kernel reads
// it, except cp, whose prefix is written and read in the same sweep
// order a fresh buffer saw, so reused scratch and fresh allocations
// solve bit-identically. dst must not alias any diagonal or d.
func TriSolve(dst, cp, dp, a, b, c, d []float64) error {
n := len(b)
if n == 0 {
return Errf("SolveTridiagonal: empty system")
}
b0 := b[0]
if b0 == 0 {
return Errf("SolveTridiagonal: zero pivot at row 0")
}
// For n = 1 the c diagonal is empty per the length contract, so the
// seed must not read it; the single unknown falls out of dp[0].
if n > 1 {
cp[0] = c[0] / b0
}
dp[0] = d[0] / b0
for i := 1; i < n; i++ {
den := b[i] - a[i-1]*cp[i-1]
if den == 0 {
return Errf("SolveTridiagonal: zero pivot at row %d", i)
}
if i < n-1 {
cp[i] = c[i] / den
}
dp[i] = (d[i] - a[i-1]*dp[i-1]) / den
}
dst[n-1] = dp[n-1]
for i := n - 2; i >= 0; i-- {
dst[i] = dp[i] - cp[i]*dst[i+1]
}
return nil
}