// Copyright (c) 2026 Petr Balvín (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 }