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

344 lines
10 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 (
"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)
}