Files
tensor/linalg/linalg.go
T

344 lines
10 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 (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Linear algebra. Every entry point requires a square
// 2-D matrix. One LU decomposition with partial pivoting backs Det,
// Solve and Inv; the same kernel serves both real and complex element
// types through generics, so complex systems solve exactly where their
// real counterparts do. A singular matrix is an error for Solve and Inv,
// while the determinants return 0. Results are float64/complex128
// approximations, as is standard for elimination with partial pivoting;
// near-singular matrices lose precision without failing.
//
// Type summary per the ladder: Solve and Inv promote to complex128 when
// any operand is complex; Det and Trace answer only real matrices and
// point complex callers at DetComplex and TraceComplex.
// squareFloatMatrix converts a real array to a square float64 matrix,
// promoting ints and float32 on the way; it errors with the caller's
// name when the shape is wrong. The rows are views over one flat
// backing slice, so the conversion costs two allocations. The rows are
// handed to the solver, which may reorder the row headers in place, so
// callers must treat the result as consumed.
func squareFloatMatrix(a *core.Array, name string) ([][]float64, error) {
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
}
n := a.Shape()[0]
back := make([]float64, n*n)
if a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == n*n {
copy(back, a.RawFloats())
} else {
for i := range n * n {
back[i] = a.FloatAt(i)
}
}
rows := make([][]float64, n)
for i := range n {
rows[i] = back[i*n : (i+1)*n]
}
return rows, nil
}
// squareComplexMatrix converts a complex array to a square matrix of
// complex128 values with the same shape contract as its real twin.
func squareComplexMatrix(a *core.Array, name string) ([][]complex128, error) {
if a.Dtype() != core.Complex {
return nil, base.Errf("%s: needs a complex matrix", name)
}
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
}
n := a.Shape()[0]
back := make([]complex128, n*n)
if !a.Strided() && len(a.RawComplexes()) == n*n {
copy(back, a.RawComplexes())
} else {
for i := range n * n {
back[i] = a.ComplexAt(i)
}
}
rows := make([][]complex128, n)
for i := range n {
rows[i] = back[i*n : (i+1)*n]
}
return rows, nil
}
// identityColumns builds the columns of the n×n identity, the standard
// right-hand side of the inverse. The columns slice into one flat
// backing array: identical values, two allocations instead of n+1.
func identityColumns[T scalar](n int) [][]T {
back := make([]T, n*n)
cols := make([][]T, n)
var one T = 1
for j := range n {
col := back[j*n : (j+1)*n]
col[j] = one
cols[j] = col
}
return cols
}
// requireReal refuses the complex dtype at an entry point that reads
// its inputs as real numbers. Without the gate the first FloatAt would
// fall through to the complex array's empty int payload and index out
// of range, which is a panic rather than the error the caller expects.
func requireReal(name string, arrays ...*core.Array) error {
for _, a := range arrays {
if a.Dtype() == core.Complex {
return base.Errf("%s: complex arrays are not supported", name)
}
}
return nil
}
// Det returns the determinant of a square real matrix (ints and
// float32 promote). A singular matrix yields 0, not an error; complex
// matrices are answered by DetComplex.
func Det(a *core.Array) (float64, error) {
if a.Dtype() == core.Complex {
return 0, base.Errf("Det: complex matrices are answered by DetComplex")
}
m, err := squareFloatMatrix(a, "Det")
if err != nil {
return 0, err
}
_, parity := base.Factor(m)
det := 1.0 * float64(parity)
for i := range m {
det *= m[i][i]
}
return det, nil
}
// DetComplex returns the determinant of a square complex matrix as a
// complex128 value. A singular matrix yields 0, not an error.
func DetComplex(a *core.Array) (complex128, error) {
m, err := squareComplexMatrix(a, "DetComplex")
if err != nil {
return 0, err
}
_, parity := base.Factor(m)
det := complex(float64(parity), 0)
for i := range m {
det *= m[i][i]
}
return det, nil
}
// solveRHS converts b into float columns of length n. b may be a vector
// of length n or a matrix with n rows.
func solveRHS(b *core.Array, n int) ([][]float64, error) {
switch {
case b.NDim() == 1 && b.Len() == n:
col := make([]float64, n)
if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n {
copy(col, b.RawFloats())
} else {
for i := range n {
col[i] = b.FloatAt(i)
}
}
return [][]float64{col}, nil
case b.NDim() == 2 && b.Shape()[0] == n:
w := b.Shape()[1]
cols := make([][]float64, w)
dense := b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n*w
for j := range w {
col := make([]float64, n)
for i := range n {
if dense {
col[i] = b.RawFloats()[i*w+j]
} else {
col[i] = b.FloatAt(i*w + j)
}
}
cols[j] = col
}
return cols, nil
}
return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s",
n, n, base.ShapeText(b.Shape()))
}
// complexRHS converts b into complex columns of length n under the same
// shape contract as solveRHS; a real right-hand side promotes wholesale.
func complexRHS(b *core.Array, n int) ([][]complex128, error) {
switch {
case b.NDim() == 1 && b.Len() == n:
col := make([]complex128, n)
if b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n {
copy(col, b.RawComplexes())
} else {
for i := range n {
col[i] = b.ComplexAt(i)
}
}
return [][]complex128{col}, nil
case b.NDim() == 2 && b.Shape()[0] == n:
w := b.Shape()[1]
cols := make([][]complex128, w)
dense := b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n*w
for j := range w {
col := make([]complex128, n)
for i := range n {
if dense {
col[i] = b.RawComplexes()[i*w+j]
} else {
col[i] = b.ComplexAt(i*w + j)
}
}
cols[j] = col
}
return cols, nil
}
return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s",
n, n, base.ShapeText(b.Shape()))
}
// columnsToFloatArray assembles solved columns back into a row-major
// array shaped like the original right-hand side: like[0] rows in both
// the vector and the matrix case.
func columnsToFloatArray(cols [][]float64, like []int) *core.Array {
rows := like[0]
out := core.New(core.Float, like...)
for i := range rows {
base := i * len(cols)
for c, col := range cols {
out.RawFloats()[base+c] = col[i]
}
}
return out
}
// columnsToComplexArray is the complex twin of columnsToFloatArray.
func columnsToComplexArray(cols [][]complex128, like []int) *core.Array {
out := core.New(core.Complex, like...)
rows := like[0]
for r := range rows {
for c, col := range cols {
out.RawComplexes()[r*len(cols)+c] = col[r]
}
}
return out
}
// Solve returns x such that a·x = b, where b is a vector of length n or
// a matrix with n rows. Real inputs give a float result; any complex
// operand promotes the whole system to complex128 per the ladder.
func Solve(a, b *core.Array) (*core.Array, error) {
if a.Dtype() == core.Complex || b.Dtype() == core.Complex {
// Any complex operand promotes the whole system: a real A
// widens alongside a complex right-hand side.
if a.Dtype() != core.Complex {
conv, cerr := core.Astype(a, core.Complex)
if cerr != nil {
return nil, cerr
}
a = conv
}
ac, err := squareComplexMatrix(a, "Solve")
if err != nil {
return nil, err
}
rhs, err := complexRHS(b, len(ac))
if err != nil {
return nil, err
}
x, err := base.SolveSystem("Solve", ac, rhs)
if err != nil {
return nil, err
}
return columnsToComplexArray(x, b.Shape()), nil
}
m, err := squareFloatMatrix(a, "Solve")
if err != nil {
return nil, err
}
rhs, err := solveRHS(b, len(m))
if err != nil {
return nil, err
}
x, err := base.SolveSystem("Solve", m, rhs)
if err != nil {
return nil, err
}
return columnsToFloatArray(x, b.Shape()), nil
}
// Inv returns the inverse of a square matrix. Real inputs give a float
// array; a complex matrix answers in complex128.
func Inv(a *core.Array) (*core.Array, error) {
if a.Dtype() == core.Complex {
ac, err := squareComplexMatrix(a, "Inv")
if err != nil {
return nil, err
}
x, err := base.SolveSystem("Inv", ac, identityColumns[complex128](len(ac)))
if err != nil {
return nil, err
}
return columnsToComplexArray(x, []int{len(ac), len(ac)}), nil
}
m, err := squareFloatMatrix(a, "Inv")
if err != nil {
return nil, err
}
x, err := base.SolveSystem("Inv", m, identityColumns[float64](len(m)))
if err != nil {
return nil, err
}
return columnsToFloatArray(x, []int{len(m), len(m)}), nil
}
// FitPolynomial fits the coefficients (lowest power first) of a
// degree-th polynomial through (x, y) samples by QR least squares over
// the Vandermonde matrix, which stays stable where the normal
// equations would square the condition number.
func FitPolynomial(x, y *core.Array, degree int) (*core.Array, error) {
n := x.Len()
if err := requireReal("FitPolynomial", x, y); err != nil {
return nil, err
}
if degree < 0 {
return nil, base.Errf("FitPolynomial: degree %d is negative", degree)
}
// degree+1 would wrap at the int maximum and slip past the sample
// count check below as a negative.
if degree == math.MaxInt {
return nil, base.Errf("FitPolynomial: degree %d is too large", degree)
}
if n < degree+1 || y.Len() != n {
return nil, base.Errf("FitPolynomial: need n ≥ degree+1 matching samples")
}
aMat, err := zeros(core.Float, []int{n, degree + 1})
if err != nil {
return nil, base.Errf("FitPolynomial: %w", err)
}
for r := range n {
pow := 1.0
for d := range degree + 1 {
aMat.SetFloatAt(r*(degree+1)+d, pow)
pow *= x.FloatAt(r)
}
}
yFlat, err := zeros(core.Float, []int{n})
if err != nil {
return nil, base.Errf("FitPolynomial: %w", err)
}
for i := range n {
yFlat.SetFloatAt(i, y.FloatAt(i))
}
return LeastSquares(aMat, yFlat)
}